{"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 gc\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom collections import Counter, defaultdict\nimport plotly.express as px\nfrom gensim.models import Word2Vec\nfrom fastai.tabular.all import *\nfrom tensorflow.keras.preprocessing.sequence import pad_sequences\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader, TensorDataset, random_split\nimport matplotlib.patches as mpatches\n\nsns.set()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-10-17T04:14:46.448993Z","iopub.execute_input":"2023-10-17T04:14:46.449775Z","iopub.status.idle":"2023-10-17T04:15:07.483378Z","shell.execute_reply.started":"2023-10-17T04:14:46.449744Z","shell.execute_reply":"2023-10-17T04:15:07.482357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"206 Reactivity columns","metadata":{}},{"cell_type":"code","source":"rna_sequence_data = pd.read_csv('/kaggle/input/stanford-ribonanza-rna-folding/train_data.csv', nrows = 25000)\nrna_sequence_data.head()","metadata":{"execution":{"iopub.status.busy":"2023-10-17T01:01:16.969133Z","iopub.execute_input":"2023-10-17T01:01:16.969819Z","iopub.status.idle":"2023-10-17T01:01:18.048338Z","shell.execute_reply.started":"2023-10-17T01:01:16.969788Z","shell.execute_reply":"2023-10-17T01:01:18.047315Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rna_sequence_data = rna_sequence_data.drop(columns=rna_sequence_data.filter(like='error').columns)\nrna_sequence_data = rna_sequence_data.drop(['sequence_id', 'dataset_name', 'reads', 'SN_filter'], axis=1)\nreactivity_cols = rna_sequence_data.filter(like='reactivity').columns\nrna_sequence_data[reactivity_cols] = rna_sequence_data[reactivity_cols].clip(lower=0)\nrna_sequence_data['reactivity'] = rna_sequence_data[reactivity_cols].mean(axis=1)\nrna_sequence_data = rna_sequence_data.drop(columns=reactivity_cols)\nrna_sequence_data","metadata":{"execution":{"iopub.status.busy":"2023-10-17T01:01:18.050267Z","iopub.execute_input":"2023-10-17T01:01:18.050844Z","iopub.status.idle":"2023-10-17T01:01:18.339600Z","shell.execute_reply.started":"2023-10-17T01:01:18.050811Z","shell.execute_reply":"2023-10-17T01:01:18.338585Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Codon Frequency Plot","metadata":{}},{"cell_type":"code","source":"overall_frequency = Counter()\n\nfor index, row in rna_sequence_data.iterrows():\n    sequence = row['sequence']\n    codons = [sequence[i:i+3] for i in range(0, len(sequence), 3)]\n    frequency = Counter(codons)\n    \n    overall_frequency.update(frequency)\n\ncustom_palette = sns.color_palette(\"husl\", len(overall_frequency))\n\nplt.figure(figsize=(16, 6))\nsns.barplot(x=list(overall_frequency.keys()), y=list(overall_frequency.values()), palette=custom_palette)\nplt.xticks(rotation=90)\nplt.xlabel('Codon')\nplt.ylabel('Frequency')\nplt.title('Overall Codon Frequency')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-17T01:01:18.341155Z","iopub.execute_input":"2023-10-17T01:01:18.342178Z","iopub.status.idle":"2023-10-17T01:01:21.067350Z","shell.execute_reply.started":"2023-10-17T01:01:18.342138Z","shell.execute_reply":"2023-10-17T01:01:21.066489Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Codon Heatmap","metadata":{}},{"cell_type":"code","source":"sequences = rna_sequence_data['sequence']\n\ncodon_freqs = []\nfor sequence in sequences:\n    codons = [sequence[i:i+3] for i in range(0, len(sequence), 3)]\n    frequency = Counter(codons)\n    codon_freqs.append(frequency)\n\ndf = pd.DataFrame(codon_freqs).fillna(0)\n\nplt.figure(figsize=(12, 8))\nsns.heatmap(df, cmap=\"YlGnBu\")\nplt.title('Heatmap of Codon Frequency')\nplt.xlabel('Codon')\n\ninterval = 10000\nplt.yticks(range(0, len(sequences), interval), range(0, len(sequences), interval))\n\nplt.ylabel('Sequence Index')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-17T01:01:21.069472Z","iopub.execute_input":"2023-10-17T01:01:21.070297Z","iopub.status.idle":"2023-10-17T01:01:23.908949Z","shell.execute_reply.started":"2023-10-17T01:01:21.070264Z","shell.execute_reply":"2023-10-17T01:01:23.908076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Nucleotide level heatmap","metadata":{}},{"cell_type":"code","source":"# nucleotides mapping\nmapping = {\n    'A': 1,\n    'C': 2,\n    'G': 3,\n    'U': 4,\n    '.': 0\n}\n\nSEQUENCE_LENGTH = 170\nnumerical_sequences = rna_sequence_data['sequence'].apply(lambda x: [mapping[n] for n in x[:SEQUENCE_LENGTH]])\nrna_map = np.array(numerical_sequences.tolist())\n\nCMAP = 'Set3'\n\nplt.figure(figsize=(16,4))\nim = plt.imshow(rna_map, aspect='auto', cmap=CMAP, vmin=0, vmax=4, interpolation='nearest')\n\nrna_values = {\n    1: 'A',\n    2: 'C',\n    3: 'G',\n    4: 'U',\n    0: 'Pad'\n}\n\ncolor_values = [1, 2, 3, 4, 0]\ncolors = [im.cmap(im.norm(color_value)) for color_value in color_values]\npatches = [mpatches.Patch(color=colors[i], label=rna_values[color_value]) for i, color_value in enumerate(color_values)]\nplt.legend(handles=patches, bbox_to_anchor=(0.65, -0.1), ncol=5, borderaxespad=0.)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-17T01:01:23.910340Z","iopub.execute_input":"2023-10-17T01:01:23.910887Z","iopub.status.idle":"2023-10-17T01:01:24.882121Z","shell.execute_reply.started":"2023-10-17T01:01:23.910855Z","shell.execute_reply":"2023-10-17T01:01:24.881200Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Composition of A,U,G,C Nucleotides","metadata":{}},{"cell_type":"code","source":"nucleotide_counts = Counter(sequence)\n\nplt.pie(nucleotide_counts.values(), labels=nucleotide_counts.keys(), autopct='%1.1f%%')\nplt.title('Nucleotide Composition Pie Chart')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-17T01:01:24.883607Z","iopub.execute_input":"2023-10-17T01:01:24.884231Z","iopub.status.idle":"2023-10-17T01:01:25.017123Z","shell.execute_reply.started":"2023-10-17T01:01:24.884200Z","shell.execute_reply":"2023-10-17T01:01:25.016162Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## GC Bond Counts on Sliding window = 30\n\n* GC bonds are more stronger than AT and AU bonds. \n* Strong GC Bonds are stable in high-temperature regions. (Organisms having more GC Bonds in DNA or RNA are more likely to survive in high temperature regions.)\n* GC bond abundant regions are more challenging to replicate leading to errors in replication.","metadata":{}},{"cell_type":"code","source":"window_size = 30 \ngc_content = [sequence[i:i+window_size].count('G') + sequence[i:i+window_size].count('C') for i in range(len(sequence) - window_size + 1)]\n\nplt.figure(figsize=(16, 6))\nplt.plot(gc_content)\nplt.xlabel('Window Start Position')\nplt.ylabel('GC Content')\nplt.title('Sliding Window GC Content')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-17T01:01:25.018802Z","iopub.execute_input":"2023-10-17T01:01:25.019463Z","iopub.status.idle":"2023-10-17T01:01:25.375182Z","shell.execute_reply.started":"2023-10-17T01:01:25.019429Z","shell.execute_reply":"2023-10-17T01:01:25.374308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Bias due to Codon Pair Interaction","metadata":{}},{"cell_type":"code","source":"overall_pair_freqs = Counter()\n\nfor seq in rna_sequence_data['sequence']:\n    pairs = [(seq[i:i+3], seq[i+3:i+6]) for i in range(0, len(seq) - 3, 3)]\n    overall_pair_freqs.update(pairs)\n\nindex = list(set([pair[0] for pair in overall_pair_freqs.keys()]))\ncolumns = list(set([pair[1] for pair in overall_pair_freqs.keys()]))\npair_df = pd.DataFrame(index=index, columns=columns).fillna(0)\n\nfor pair, freq in overall_pair_freqs.items():\n    pair_df.at[pair[0], pair[1]] = freq\n    \nplt.figure(figsize=(10, 6))\nsns.heatmap(pair_df, cmap='viridis')\nplt.title('Overall Codon Pair Bias')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-17T01:01:25.376806Z","iopub.execute_input":"2023-10-17T01:01:25.377460Z","iopub.status.idle":"2023-10-17T01:01:26.826651Z","shell.execute_reply.started":"2023-10-17T01:01:25.377425Z","shell.execute_reply":"2023-10-17T01:01:26.825695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 3D Scatter Plot of Codon Counts for AUG, GCA, and GCU Codons","metadata":{}},{"cell_type":"code","source":"codon_counts = defaultdict(list)\n\nfor seq in rna_sequence_data['sequence']:\n    counter = Counter([seq[i:i+3] for i in range(0, len(seq), 3)])\n    \n    for codon in counter:\n        codon_counts[codon].append(counter[codon])\n\nfor codon, counts in codon_counts.items():\n    while len(counts) < len(rna_sequence_data['sequence']):\n        counts.append(0)\n\ncodon_count = pd.DataFrame(dict(codon_counts))\n\nfig = px.scatter_3d(codon_count, x='AUG', y='GCA', z='GCU')\nfig.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-17T01:01:26.827959Z","iopub.execute_input":"2023-10-17T01:01:26.828579Z","iopub.status.idle":"2023-10-17T01:01:30.004171Z","shell.execute_reply.started":"2023-10-17T01:01:26.828541Z","shell.execute_reply":"2023-10-17T01:01:30.003308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create Embeddings (Word2Vec)","metadata":{}},{"cell_type":"code","source":"def sequence_to_kmers(sequence, k=3):\n    \"\"\"Convert sequence to overlapping k-mers.\"\"\"\n    return [sequence[i:i+k] for i in range(len(sequence) - k + 1)]\n\nkmers = rna_sequence_data['sequence'].apply(sequence_to_kmers)\n\nmodel = Word2Vec(sentences=kmers, vector_size=100, window=5, min_count=1, workers=4)\n\ndef sequence_to_vector(sequence, model):\n    kmers = sequence_to_kmers(sequence)\n    vectors = [model.wv[kmer] for kmer in kmers if kmer in model.wv]\n    return sum(vectors) / len(vectors)","metadata":{"execution":{"iopub.status.busy":"2023-10-17T01:01:30.007794Z","iopub.execute_input":"2023-10-17T01:01:30.008414Z","iopub.status.idle":"2023-10-17T01:01:41.686346Z","shell.execute_reply.started":"2023-10-17T01:01:30.008379Z","shell.execute_reply":"2023-10-17T01:01:41.685463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rna_sequence_data['vectors'] = rna_sequence_data['sequence'].apply(lambda x: sequence_to_vector(x, model))","metadata":{"execution":{"iopub.status.busy":"2023-10-17T01:01:41.687985Z","iopub.execute_input":"2023-10-17T01:01:41.688407Z","iopub.status.idle":"2023-10-17T01:01:50.214276Z","shell.execute_reply.started":"2023-10-17T01:01:41.688374Z","shell.execute_reply":"2023-10-17T01:01:50.213206Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rna_sequence_data = rna_sequence_data.dropna(subset=['reactivity'])\n\ncols_to_drop = ['sequence', 'experiment_type', 'signal_to_noise']\nexperiment_2A3_MaP = rna_sequence_data[rna_sequence_data['experiment_type'] == '2A3_MaP'].drop(columns=cols_to_drop).reset_index(drop=True)\nexperiment_DMS_MaP = rna_sequence_data[rna_sequence_data['experiment_type'] == 'DMS_MaP'].drop(columns=cols_to_drop).reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2023-10-17T01:01:50.215728Z","iopub.execute_input":"2023-10-17T01:01:50.216818Z","iopub.status.idle":"2023-10-17T01:01:50.235987Z","shell.execute_reply.started":"2023-10-17T01:01:50.216782Z","shell.execute_reply":"2023-10-17T01:01:50.235178Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"experiment_DMS_MaP.head()","metadata":{"execution":{"iopub.status.busy":"2023-10-17T01:01:50.237154Z","iopub.execute_input":"2023-10-17T01:01:50.237473Z","iopub.status.idle":"2023-10-17T01:01:50.250462Z","shell.execute_reply.started":"2023-10-17T01:01:50.237444Z","shell.execute_reply":"2023-10-17T01:01:50.249587Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Approaching it with average Reactivity","metadata":{}},{"cell_type":"code","source":"vector_data = pd.DataFrame(experiment_2A3_MaP['vectors'].to_list(), columns=[f'feature_{i}' for i in range(100)])\nexp_2a3 = pd.concat([experiment_2A3_MaP, vector_data], axis=1).drop(columns='vectors')\n\ncont_names = [f'feature_{i}' for i in range(100)]\ndep_var = 'reactivity'\n\nsplits = RandomSplitter(valid_pct=0.3)(range_of(exp_2a3))  # 70% train, 30% validation\nto = TabularPandas(exp_2a3, procs=[Normalize], cont_names=cont_names, y_names=dep_var, splits=splits)\n\ndls = to.dataloaders(bs=64)\nlearn = tabular_learner(dls, layers=[300,200], metrics=mae)\n\nlearn.lr_find(start_lr=1e-7, end_lr=10, num_it=100)\n\nlearn.fit_one_cycle(1, 1e-2)\n\nlearn.recorder.plot_lr_find()\n\nlearn.recorder.plot_loss()","metadata":{"execution":{"iopub.status.busy":"2023-10-17T01:01:50.252050Z","iopub.execute_input":"2023-10-17T01:01:50.252416Z","iopub.status.idle":"2023-10-17T01:01:54.641340Z","shell.execute_reply.started":"2023-10-17T01:01:50.252383Z","shell.execute_reply":"2023-10-17T01:01:54.640303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del vector_data\ndel exp_2a3\ndel experiment_DMS_MaP\ndel experiment_2A3_MaP\ndel df\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-10-17T01:01:54.642745Z","iopub.execute_input":"2023-10-17T01:01:54.643299Z","iopub.status.idle":"2023-10-17T01:01:58.229547Z","shell.execute_reply.started":"2023-10-17T01:01:54.643267Z","shell.execute_reply":"2023-10-17T01:01:58.228631Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# It's a Sequence to Sequence Problem, Let's approach it using Seq-to-Seq RNN","metadata":{}},{"cell_type":"markdown","source":"# Data Preprocessing","metadata":{}},{"cell_type":"code","source":"rna_sequence_data = pd.read_csv('/kaggle/input/stanford-ribonanza-rna-folding/train_data.csv', nrows = 150000)\nrna_sequence_data.head()","metadata":{"execution":{"iopub.status.busy":"2023-10-17T04:15:07.485063Z","iopub.execute_input":"2023-10-17T04:15:07.486040Z","iopub.status.idle":"2023-10-17T04:15:18.998078Z","shell.execute_reply.started":"2023-10-17T04:15:07.486006Z","shell.execute_reply":"2023-10-17T04:15:18.997105Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rna_sequence_data = rna_sequence_data.drop(columns=rna_sequence_data.filter(like='error').columns)\nrna_sequence_data = rna_sequence_data.drop(['sequence_id', 'dataset_name', 'reads', 'SN_filter'], axis=1)\nreactivity_cols = rna_sequence_data.filter(like='reactivity').columns\nrna_sequence_data[reactivity_cols] = rna_sequence_data[reactivity_cols].clip(lower=0)\nrow_means = rna_sequence_data[reactivity_cols].mean(axis=1)\nrna_sequence_data[reactivity_cols] = rna_sequence_data[reactivity_cols].apply(lambda col: col.fillna(row_means))\nrna_sequence_data[reactivity_cols] = rna_sequence_data[reactivity_cols].fillna(0)\n\ndef encode_nucleotide(nucleotide):\n    mapping = {'A': [1, 0, 0, 0],\n               'U': [0, 1, 0, 0],\n               'G': [0, 0, 1, 0],\n               'C': [0, 0, 0, 1]}\n    return mapping[nucleotide]\n\ndef encode_sequence(sequence):\n    return [encode_nucleotide(n) for n in sequence]\n\nrna_sequence_data['encoded_sequence'] = rna_sequence_data['sequence'].apply(encode_sequence)\n\nrna_sequence_data['padded_sequence'] = rna_sequence_data['encoded_sequence'].apply(lambda x: pad_sequences([x], maxlen=457, padding='post')[0])\n\ndef pad_reactivity_sequence(sequence, max_length=457, padding_value=-0.1):\n    # Calculate how many padding values are needed\n    padding_length = max_length - len(sequence)\n    \n    # Create the padded sequence\n    padded_sequence = list(sequence) + [padding_value] * padding_length\n    \n    return padded_sequence\n\n# 1. Consolidate Reactivity Columns\nreactivity_columns = [f'reactivity_{i:04}' for i in range(1, 207)]  # Adjust this based on the number of columns you have\n\nrna_sequence_data['reactivity'] = rna_sequence_data[reactivity_columns].values.tolist()\n\n# Drop the individual reactivity columns as they're now redundant (optional)\nrna_sequence_data.drop(columns=reactivity_columns, inplace=True)\n\n# Apply the padding function\nrna_sequence_data['padded_reactivity'] = rna_sequence_data['reactivity'].apply(pad_reactivity_sequence)\n\n# 1. Store original sequence lengths\nrna_sequence_data['original_length'] = rna_sequence_data['reactivity'].apply(len)\n\n# 2. Inverse transform function\ndef inverse_transform(padded_sequence, original_length):\n    return padded_sequence[:original_length]\n\n# Apply the inverse transform\nrna_sequence_data['retrieved_reactivity'] = rna_sequence_data.apply(lambda row: inverse_transform(row['padded_reactivity'], row['original_length']), axis=1)","metadata":{"execution":{"iopub.status.busy":"2023-10-17T04:15:23.043585Z","iopub.execute_input":"2023-10-17T04:15:23.044038Z","iopub.status.idle":"2023-10-17T04:17:11.277061Z","shell.execute_reply.started":"2023-10-17T04:15:23.044004Z","shell.execute_reply":"2023-10-17T04:17:11.276055Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rna_2A3 = rna_sequence_data[rna_sequence_data['experiment_type'] == '2A3_MaP']\nrna_dms = rna_sequence_data[rna_sequence_data['experiment_type'] == 'DMS_MaP']\n\ndel rna_sequence_data\ngc.collect()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\nclass RNASeq2SeqModel(nn.Module):\n    def __init__(self, embedding_dim, hidden_dim, num_layers):\n        super(RNASeq2SeqModel, self).__init__()\n        \n        # Input is 4 due to one-hot encoding of A, U, G, C\n        self.embedding = nn.Embedding(4, embedding_dim)  # Embedding layer\n        self.rnn = nn.LSTM(embedding_dim, hidden_dim, num_layers, batch_first=True)\n        self.fc = nn.Linear(hidden_dim, 1)  # Predicting reactivity for each nucleotide\n        \n    def forward(self, x):\n        x = torch.argmax(x, dim=-1)  # Convert one-hot to indices for embedding\n        x = self.embedding(x)\n        rnn_out, _ = self.rnn(x)\n        output = self.fc(rnn_out)\n        return output\n\n    \nsequence_matrix = rna_2A3['padded_sequence']\nsequence_tensor = torch.FloatTensor(sequence_matrix)\nreactivities = torch.FloatTensor(rna_2A3['padded_reactivity'])\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training RNN Seq-to-Seq on 2A3_MaP data","metadata":{}},{"cell_type":"code","source":"# Create Dataset\ndataset = TensorDataset(sequence_tensor, reactivities)\n\n# Split data into train and validation sets (80% train, 20% validation for this example)\ntrain_size = int(0.8 * len(dataset))\nval_size = len(dataset) - train_size\ntrain_dataset, val_dataset = random_split(dataset, [train_size, val_size])\n\ntrain_loader = DataLoader(train_dataset, batch_size=64, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=64, shuffle=False)\n\n# Hyperparameters\nembedding_dim = 16\nhidden_dim = 32\nnum_layers = 2\nlearning_rate = 0.001\nnum_epochs = 100\n\n# Instantiate model\nmodel_2a3 = RNASeq2SeqModel(embedding_dim=16, hidden_dim=32, num_layers=2).to(device)\ncriterion = nn.MSELoss()\noptimizer = torch.optim.Adam(model_2a3.parameters(), lr=0.001)\n\nfor epoch in range(num_epochs):\n    \n    model_2a3.train()\n    total_loss = 0.0\n    \n    for batch_seq, batch_react in train_loader:\n        \n        batch_seq, batch_react = batch_seq.to(device), batch_react.to(device) \n        \n        # Zero the parameter gradients\n        optimizer.zero_grad()\n        \n        # Forward pass\n        outputs = model_2a3(batch_seq)\n        \n        # Compute loss\n        loss = criterion(outputs.squeeze(), batch_react)\n        \n        # Backward pass and optimize\n        loss.backward()\n        optimizer.step()\n        \n        total_loss += loss.item()\n    \n    # Validate and compute MSE for the validation set\n    model_2a3.eval()\n    val_loss = 0.0\n    with torch.no_grad():\n        for batch_seq, batch_react in val_loader:\n            batch_seq, batch_react = batch_seq.to(device), batch_react.to(device)  # Transfer data to GPU\n            outputs = model_2a3(batch_seq)\n            loss = criterion(outputs.squeeze(), batch_react)\n            val_loss += loss.item()\n        \n    print(f\"Epoch {epoch+1}/{num_epochs}, Training Loss: {total_loss/len(train_loader)}, Validation MSE: {val_loss/len(val_loader)}\")\n\nprint(\"Training finished.\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del dataset\ndel train_dataset\ndel val_dataset\ndel train_loader\ndel val_loader\ndel sequence_tensor\ndel reactivities\n\ngc.collect()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training RNN Seq-to-Seq on DMS_MaP data","metadata":{}},{"cell_type":"code","source":"sequence_matrix_dms = np.stack(rna_dms['padded_sequence'].to_numpy())  # Convert Series of lists to a numpy array\nsequence_tensor_dms = torch.FloatTensor(sequence_matrix_dms)\nreactivities_dms = torch.FloatTensor(np.stack(rna_dms['padded_reactivity'].to_numpy()))\n\n# Create Dataset for rna_dms\ndataset_dms = TensorDataset(sequence_tensor_dms, reactivities_dms)\n\n# Split data into train and validation sets (80% train, 20% validation for this example)\ntrain_size_dms = int(0.8 * len(dataset_dms))\nval_size_dms = len(dataset_dms) - train_size_dms\ntrain_dataset_dms, val_dataset_dms = random_split(dataset_dms, [train_size_dms, val_size_dms])\n\ntrain_loader_dms = DataLoader(train_dataset_dms, batch_size=64, shuffle=True)\nval_loader_dms = DataLoader(val_dataset_dms, batch_size=64, shuffle=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Instantiate another model for rna_dms\nmodel_dms = RNASeq2SeqModel(embedding_dim=embedding_dim, hidden_dim=hidden_dim, num_layers=num_layers).to(device)\noptimizer_dms = torch.optim.Adam(model_dms.parameters(), lr=learning_rate)\n\n# Training Loop for rna_dms\nfor epoch in range(num_epochs):\n    model_dms.train()\n    total_loss = 0.0\n    for batch_seq, batch_react in train_loader_dms:\n        batch_seq, batch_react = batch_seq.to(device), batch_react.to(device) \n        optimizer_dms.zero_grad()\n        \n        # Forward pass\n        outputs = model_dms(batch_seq)\n        \n        # Compute loss\n        loss = criterion(outputs.squeeze(), batch_react)\n        \n        # Backward pass and optimize\n        loss.backward()\n        optimizer_dms.step()\n        \n        total_loss += loss.item()\n    \n    # Validate and compute MSE for the validation set\n    model_dms.eval()\n    val_loss = 0.0\n    with torch.no_grad():\n        for batch_seq, batch_react in val_loader_dms:\n            batch_seq, batch_react = batch_seq.to(device), batch_react.to(device)  # Transfer data to GPU\n            outputs = model_dms(batch_seq)\n            loss = criterion(outputs.squeeze(), batch_react)\n            val_loss += loss.item()\n        \n    \n    print(f\"[DMS] Epoch {epoch+1}/{num_epochs}, Training Loss: {total_loss/len(train_loader_dms)}, Validation MSE: {val_loss/len(val_loader_dms)}\")\n\nprint(\"Training for DMS finished.\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model_dms, \"/kaggle/working/models/model_dms.pth\")\ntorch.save(model_2a3, \"/kaggle/working/models/model_2a3.pth\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}