{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":87793,"databundleVersionId":11553390,"sourceType":"competition"}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 🧬 RNA-FM Embeddings: Enabling Efficient Similarity Search with k-NN\n\n## 📌 Introduction  \nThis notebook demonstrates how to extract meaningful vector embeddings from RNA sequences using the **RNA-FM** foundation model. https://arxiv.org/pdf/2204.00300\n\nI decided to use **RNA-FM** as the pretrained model because it was trained for multiple classification, regression, and structure prediction objectives, so it will have some context of the RNA shape.\n\n---\n\n## 🚀 Why This Approach Is Useful  \nIf we use the raw output from the models, we won't get good results for 3D structure prediction, but we can also use the embeddings to find similar sequences which could be helpful for creating a model that generates new 3D structures based on the most similar existing ones in the dataset. This kind of approach has shown some utility for other protein-related problems.\n\n\n","metadata":{}},{"cell_type":"code","source":"!pip install multimolecule -q","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T06:08:47.368322Z","iopub.execute_input":"2025-04-01T06:08:47.368611Z","iopub.status.idle":"2025-04-01T06:08:51.072269Z","shell.execute_reply.started":"2025-04-01T06:08:47.368587Z","shell.execute_reply":"2025-04-01T06:08:51.071228Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np \nimport torch\nfrom tqdm import tqdm\nimport plotly.graph_objects as go\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T06:08:57.901382Z","iopub.execute_input":"2025-04-01T06:08:57.901679Z","iopub.status.idle":"2025-04-01T06:09:02.376348Z","shell.execute_reply.started":"2025-04-01T06:08:57.901654Z","shell.execute_reply":"2025-04-01T06:09:02.375645Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Extraction of base Dataframes\ntrain_labels = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/train_labels.csv')\ntest_labels = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/validation_labels.csv')\ntrain_df = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/train_sequences.csv')\ntest_df = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/test_sequences.csv')\n\n# Extraction of RNA sequences from DataFrames\ntrain_rna, test_rna = train_df['sequence'].to_list(), test_df['sequence'].to_list()\n\n# Create a unique identifier by extracting the base ID from the full ID\n# (removes the last part after the last underscore)\ntrain_labels['unique_id'] = train_labels['ID'].str.rsplit('_', n=1).str[0]\ntest_labels['unique_id'] = test_labels['ID'].str.rsplit('_', n=1).str[0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T06:11:43.006995Z","iopub.execute_input":"2025-04-01T06:11:43.007306Z","iopub.status.idle":"2025-04-01T06:11:43.039821Z","shell.execute_reply.started":"2025-04-01T06:11:43.007281Z","shell.execute_reply":"2025-04-01T06:11:43.039175Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from multimolecule import RnaTokenizer, RnaFmModel\n\n# Define constants\nMAX_SEQ_LEN = 720\nembed_dim = 640  # Size of the RNA-FM output tensor\nbatch_size = 32\n\n# Make sure device is defined\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Load tokenizer and model\ntokenizer = RnaTokenizer.from_pretrained(\"multimolecule/rnafm\")\nmodel = RnaFmModel.from_pretrained(\"multimolecule/rnafm\").to(device)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T06:09:10.278260Z","iopub.execute_input":"2025-04-01T06:09:10.278686Z","iopub.status.idle":"2025-04-01T06:09:32.087524Z","shell.execute_reply.started":"2025-04-01T06:09:10.278660Z","shell.execute_reply":"2025-04-01T06:09:32.086846Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Embedding Generation \n<hr>","metadata":{}},{"cell_type":"code","source":"def extract_embeddings(rna_seq, batch_size=32, MAX_SEQ_LEN=720):\n    \"\"\"\n    Extracts embeddings from RNA sequences using a pre-trained language model.\n    \n    Parameters:\n    - rna_seq: List of RNA sequences (strings of nucleotides)\n    - batch_size: Number of sequences to process at once (default: 32)\n    - MAX_SEQ_LEN: Maximum sequence length for tokenization (default: 720)\n    \n    Returns:\n    - last_hidden_states_stack: NumPy array of shape (n_sequences, sequence_length, hidden_size)\n      containing contextual embeddings for each token in each sequence\n    - pooler_outputs_stack: NumPy array of shape (n_sequences, hidden_size)\n      containing aggregated sequence-level embeddings\n    \"\"\"\n    # To process multiple sequences from a DataFrame\n    last_hidden_states = []\n    pooler_outputs = []\n    \n    # Using batch processing for faster inference\n    for idx in tqdm(range(0, len(rna_seq), batch_size)):\n    \n        seq = rna_seq[idx:idx+batch_size]\n        \n        # Tokenize and get model outputs\n        inputs = tokenizer(\n            seq, \n            return_tensors=\"pt\",\n            padding=\"max_length\",  # Add padding\n            truncation=True,       # Enable truncation\n            max_length=MAX_SEQ_LEN # Set the maximum length that the model expects\n        ).to(device)    \n        \n        with torch.no_grad():\n            outputs = model(**inputs)\n        \n        # Extract and store tensors as NumPy arrays\n        last_hidden_states.extend(outputs.last_hidden_state.detach().cpu().numpy())\n        pooler_outputs.extend(outputs.pooler_output.detach().cpu().numpy())\n    \n    # Convert lists to NumPy arrays\n    last_hidden_states_stack = np.array(last_hidden_states)\n    pooler_outputs_stack = np.array(pooler_outputs)\n    \n    return last_hidden_states_stack, pooler_outputs_stack\n    \n\n# Use of the function for train and test sequences\nhidden_states_train, pooler_outputs_train = extract_embeddings(train_rna, batch_size=16)\nhidden_states_test, pooler_outputs_test = extract_embeddings(test_rna, batch_size=16)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T06:21:23.652728Z","iopub.execute_input":"2025-04-01T06:21:23.653220Z","iopub.status.idle":"2025-04-01T06:21:54.845674Z","shell.execute_reply.started":"2025-04-01T06:21:23.653181Z","shell.execute_reply":"2025-04-01T06:21:54.844754Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.neighbors import NearestNeighbors\n\ndef k_similars_embeddings(base_embedding, embeddings_list, k=5):\n    \"\"\"\n    Find k most similar embeddings to the base embedding using cosine similarity.\n    \n    Args:\n        base_embedding: The reference embedding to compare against\n        embeddings_list: List of embeddings to search through\n        k: Number of similar embeddings to return (default: 5)\n    \n    Returns:\n        indices: Indices of the k most similar embeddings\n        distances: Corresponding similarity distances\n    \"\"\"\n\n    # Reshape base_embedding if it's only one sample\n    if base_embedding.shape[0] == 640:\n        base_embedding = base_embedding.reshape(1, -1)\n\n    \n    # Cosine similarity is prefered for embeddings as it measures angular distance\n    knn = NearestNeighbors(n_neighbors=k, algorithm='auto', metric='cosine')\n    knn.fit(embeddings_list)\n    \n    # Find the k most similar embeddings and their distances\n    distances, indices = knn.kneighbors(base_embedding)\n    return indices, distances\n\n\n# Finding 5 most similar proteins found in train data for the test sequences\ntest_indices, test_distances = k_similars_embeddings(pooler_outputs_test, pooler_outputs_train, k=10)\ntrain_indices, train_distances = k_similars_embeddings(pooler_outputs_train, pooler_outputs_train, k=10)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T08:13:39.862176Z","iopub.execute_input":"2025-04-01T08:13:39.862457Z","iopub.status.idle":"2025-04-01T08:13:39.890211Z","shell.execute_reply.started":"2025-04-01T08:13:39.862435Z","shell.execute_reply":"2025-04-01T08:13:39.889426Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##  Plotting sequences with similar embeddings\n<hr>","metadata":{}},{"cell_type":"code","source":"import plotly.graph_objects as go\n\n\ndef plot_structure(similar_idx, distances, n_similars=3, max_sizes=720) -> None:\n    \"\"\"\n    Plot 3D structures of RNA molecules with similar sequences.\n    \n    Args:\n        similar_idx: Indices of similar RNA sequences\n        distances: Distance metrics between sequences\n        n_similars: Number of similar structures to display (default: 3)\n        max_sizes: Maximum size parameter (default: 720)\n    \"\"\"\n    # Limit the number of structures to display\n    similar_idx, distances = similar_idx[:n_similars], distances[:n_similars]\n    \n    # Get base IDs, sequences and labels from the training data\n    base_ids = train_df.iloc[similar_idx]['target_id'].values\n    sequences = train_df.iloc[similar_idx]['sequence'].values\n    similar_df_labels = train_labels[train_labels['unique_id'].isin(base_ids)]\n    \n    # Create a dictionary mapping sequences to their 3D coordinates\n    coordinates = {}\n    for curr_id, data in similar_df_labels.groupby('unique_id'):\n        coordinates[train_df[train_df['target_id'] == curr_id]['sequence'].values[0]] = data[['x_1', 'y_1', 'z_1']].values\n    \n    # Define colors for nucleotides\n    nucleotide_colors = {\"A\": \"red\", \"G\": \"blue\", \"C\": \"green\", \"U\": \"orange\"}\n    \n    # Define colors for each sequence backbone\n    backbone_colors = [\"rgba(255,0,0,0.7)\", \"rgba(0,0,255,0.7)\", \"rgba(0,255,0,0.7)\", \n                      \"rgba(255,165,0,0.7)\", \"rgba(128,0,128,0.7)\", \"rgba(0,128,128,0.7)\"]\n    \n    fig = go.Figure()\n    \n    # Preprocess coordinates to center them\n    processed_coordinates = {}\n    centroids = []\n    \n    # First, center each structure on its own centroid\n    for i, sequence in enumerate(sequences):\n        if sequence in coordinates:\n            x, y, z = coordinates[sequence][:, 0], coordinates[sequence][:, 1], coordinates[sequence][:, 2]\n            \n            # Check that lists have the same length\n            if not (len(x) == len(y) == len(z) == len(sequence)):\n                print(f\"Warning: Lists for sequence {i+1} don't have the same length. Adjusting...\")\n                min_len = min(len(x), len(y), len(z), len(sequence))\n                x, y, z = x[:min_len], y[:min_len], z[:min_len]\n                sequence = sequence[:min_len]\n            \n            # Calculate the centroid of this structure\n            centroid = np.array([np.mean(x), np.mean(y), np.mean(z)])\n            centroids.append(centroid)\n            \n            # Center the coordinates\n            centered_coords = np.column_stack((x, y, z)) - centroid\n            processed_coordinates[sequence] = centered_coords\n    \n    # Calculate small displacements for each structure\n    max_radius = max([np.max(np.sqrt(np.sum(coords**2, axis=1))) for coords in processed_coordinates.values()])\n    spacing = max_radius * 0.5  # Space between structures\n    \n    # Iterate over all available sequences\n    for i, sequence in enumerate(sequences):\n        if sequence in coordinates:\n            sizes = len(sequence)\n            \n            # Get centered coordinates\n            centered_coords = processed_coordinates[sequence]\n            \n            # Apply a small radial displacement to visualize structures together\n            # but not completely superimposed\n            angle = 2 * np.pi * i / len(sequences)  # Distribute in circle\n            offset = np.array([spacing * np.cos(angle), spacing * np.sin(angle), 0])\n            \n            x = centered_coords[:, 0] + offset[0]\n            y = centered_coords[:, 1] + offset[1]\n            z = centered_coords[:, 2] + offset[2]\n            \n            # Select color for the backbone (rotating if there are more sequences than colors)\n            backbone_color = backbone_colors[i % len(backbone_colors)]\n            \n            # Add points by nucleotide type for this sequence\n            for resname, color in nucleotide_colors.items():\n                indices = [j for j, res in enumerate(sequence) if res == resname]\n                if indices:\n                    fig.add_trace(go.Scatter3d(\n                        x=[x[j] for j in indices],\n                        y=[y[j] for j in indices],\n                        z=[z[j] for j in indices],\n                        mode='markers',\n                        marker=dict(size=4, color=color),\n                        name=f'{resname} (Seq {i+1})',\n                        # Only show in legend for the first sequence to avoid duplicates\n                        showlegend=(i == 0)\n                    ))\n                    \n            # Add line for the RNA backbone\n            fig.add_trace(go.Scatter3d(\n                x=x,\n                y=y,\n                z=z,\n                mode='lines',\n                line=dict(color=backbone_color, width=8),\n                name=f'RNA Backbone (Seq {i+1} Similarity:{distances[i]*1000:.4f})'\n            ))\n    \n    fig.update_layout(\n        scene=dict(\n            xaxis_title='X',\n            yaxis_title='Y',\n            zaxis_title='Z',\n            aspectmode='data'\n        ),\n        title=f'Comparison of RNA 3D structures ({len(sequences)} sequences)',\n        legend=dict(\n            itemsizing='constant',\n            itemwidth=30\n        ),\n    )\n    \n    # Display the figure\n    fig.show()\n\n\n# plot sequences from most similar embeddings for train sample 0 \nplot_structure(train_indices[0], train_distances[0], n_similars=3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T08:20:51.028383Z","iopub.execute_input":"2025-04-01T08:20:51.028694Z","iopub.status.idle":"2025-04-01T08:20:51.071184Z","shell.execute_reply.started":"2025-04-01T08:20:51.028669Z","shell.execute_reply":"2025-04-01T08:20:51.070485Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# plot sequences from most similar embeddings for test sample 8\nplot_structure(test_indices[8], test_distances[8], n_similars=2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T08:20:00.507633Z","iopub.execute_input":"2025-04-01T08:20:00.507975Z","iopub.status.idle":"2025-04-01T08:20:00.538050Z","shell.execute_reply.started":"2025-04-01T08:20:00.507947Z","shell.execute_reply":"2025-04-01T08:20:00.537346Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}