{"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":"none","dataSources":[{"sourceId":87793,"databundleVersionId":12024591,"sourceType":"competition"},{"sourceId":7639698,"sourceType":"datasetVersion","datasetId":4299272},{"sourceId":8318191,"sourceType":"datasetVersion","datasetId":4459124},{"sourceId":224703571,"sourceType":"kernelVersion"}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# **RNA 3D Structure Prediction - Dataset Overview & Concepts**\n\n## **1. Introduction to RNA Structure**\nRibonucleic Acid (RNA) plays a crucial role in biological functions, including gene expression, regulation, and enzymatic activities. Unlike DNA, RNA is single-stranded and can fold into complex **secondary** and **tertiary** structures, influencing its function.\n\n### **RNA Structural Hierarchy**\n- **Primary Structure**: The linear sequence of nucleotides (A, C, G, U).\n- **Secondary Structure**: Base-pairing interactions (e.g., helices, loops).\n- **Tertiary Structure**: The full 3D shape formed by folding secondary structures.\n\nPredicting **RNA’s 3D structure** from its **sequence** remains a challenge due to:\n1. The flexibility and variability of RNA folding.\n2. Limited experimental 3D structures in databases.\n3. The need for advanced modeling approaches, including **deep learning**.\n\n---\n\n## **2. Dataset Description**\nThis competition provides RNA sequences and their corresponding **C1' atom coordinates** (a key structural reference in RNA molecules). The dataset includes:\n\n### **Train, Validation, and Test Sequences**\n- **[train/validation/test]_sequences.csv**:\n  - `target_id`: Unique identifier (e.g., `pdb_id_chain_id` in train).\n  - `sequence`: The **RNA sequence** consisting of **A, C, G, U**.\n  - `temporal_cutoff`: The date when the sequence was published.\n  - `description`: Biological metadata.\n  - `all_sequences`: FASTA-formatted full molecular sequences.\n\n### **3D Structural Labels (Ground Truth)**\n- **[train/validation]_labels.csv**:\n  - `ID`: Unique identifier (`target_id_residue_number`).\n  - `resname`: RNA nucleotide (A, C, G, U).\n  - `resid`: Residue number.\n  - `x_1, y_1, z_1, ...`: **C1' atom coordinates** in **Angstroms**.\n  - Some sequences have **multiple conformations**, recorded as `x_2, y_2, z_2`, etc.\n\n### **Submission File Format**\n- **sample_submission.csv** follows the same structure as `train_labels.csv` but requires **five sets of predicted 3D coordinates** per RNA residue.\n\n---\n\n## **3. Multiple Sequence Alignments (MSA)**\n- **MSA files (`{target_id}.MSA.fasta`)** provide evolutionary context by aligning similar RNA sequences.\n- MSA data helps **identify conserved structural elements** for better predictions.\n\n---\n\n\n## **4. External Data & Pretrained Models**\n- **RibonanzaNet**: A deep learning foundation model for RNA structures.\n- **RFdiffusion Data**: A synthetic dataset of **400,000 RNA structures** for additional training.\n- **PDB Database**: Can be used after leaderboard reset.\n\n\n","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\n# Load train sequences\ntrain_sequences = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_sequences.csv\")\n\n# Load train labels (3D coordinates)\ntrain_labels = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_labels.csv\")\n\n# Load validation sequences\nvalidation_sequences = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/validation_sequences.csv\")\n\n# Load validation labels (3D coordinates)\nvalidation_labels = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/validation_labels.csv\")\n\n# Load test sequences\ntest_sequences = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/test_sequences.csv\")\n\n# Display basic info\nprint(\"Train Sequences Sample:\")\nprint(train_sequences.head())\n\nprint(\"\\nTrain Labels Sample:\")\nprint(train_labels.head())\n\nprint(\"\\nValidation Sequences Sample:\")\nprint(validation_sequences.head())\n\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-01T16:43:22.585456Z","iopub.execute_input":"2025-05-01T16:43:22.585783Z","iopub.status.idle":"2025-05-01T16:43:24.334305Z","shell.execute_reply.started":"2025-05-01T16:43:22.585747Z","shell.execute_reply":"2025-05-01T16:43:24.333289Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\nfrom mpl_toolkits.mplot3d import Axes3D\n\n# -----------------------------\n# 1️⃣ Load and Preprocess Data\n# -----------------------------\n\n# Load train_labels.csv\ntrain_labels = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_labels.csv\")  # Replace with actual path\n\n# Extract `target_id` from `ID` (Remove residue numbering)\ntrain_labels[\"target_id\"] = train_labels[\"ID\"].apply(lambda x: \"_\".join(x.split(\"_\")[:2]))\n\n# Sort data by `target_id` and `resid` for proper ordering\ntrain_labels = train_labels.sort_values(by=[\"target_id\", \"resid\"])\n\n# -----------------------------\n# 2️⃣ Select 3 Sample RNA Sequences for Visualization\n# -----------------------------\n\n# Select first 3 unique target_ids\nsample_rnas = train_labels[\"target_id\"].unique()[:9]\n\n# Color mapping for nucleotides (A, C, G, U)\nnucleotide_colors = {\"A\": \"red\", \"C\": \"blue\", \"G\": \"green\", \"U\": \"purple\"}\n\n# -----------------------------\n# 3️⃣ 3D Plotting Function\n# -----------------------------\n\ndef plot_rna_structure_with_edges(ax, rna_data, target_id):\n    \"\"\"Plots the 3D structure of an RNA sequence with edges between consecutive nucleotides.\"\"\"\n    for resname, color in nucleotide_colors.items():\n        subset = rna_data[rna_data[\"resname\"] == resname]\n        ax.scatter(subset[\"x_1\"], subset[\"y_1\"], subset[\"z_1\"], c=color, label=resname, s=30)\n        \n    # Draw edges between consecutive nucleotides\n    for i in range(len(rna_data) - 1):\n        x_vals = [rna_data.iloc[i][\"x_1\"], rna_data.iloc[i+1][\"x_1\"]]\n        y_vals = [rna_data.iloc[i][\"y_1\"], rna_data.iloc[i+1][\"y_1\"]]\n        z_vals = [rna_data.iloc[i][\"z_1\"], rna_data.iloc[i+1][\"z_1\"]]\n        ax.plot(x_vals, y_vals, z_vals, color=\"gray\", linewidth=0.5)  # Edge line between consecutive residues\n\n    ax.set_xlabel(\"X Coordinate (Å)\")\n    ax.set_ylabel(\"Y Coordinate (Å)\")\n    ax.set_zlabel(\"Z Coordinate (Å)\")\n    ax.set_title(f\"3D Structure of RNA: {target_id}\")\n    ax.legend()\n\n# -----------------------------\n# 4️⃣ Plot 3D Structures with Edges for Selected RNAs\n# -----------------------------\n\n\n# Define the number of rows and columns for the grid\nnum_rows, num_cols = 5, 5\n\n# Select the first 9 unique target_ids for visualization\nsample_rnas = train_labels[\"target_id\"].unique()[: num_rows * num_cols]\n\n# Create a figure for 9 subplots (3x3 layout)\nfig = plt.figure(figsize=(15, 15))\n\nfor i, target_id in enumerate(sample_rnas):\n    ax = fig.add_subplot(num_rows, num_cols, i + 1, projection=\"3d\")\n    \n    # Filter data for the selected RNA sequence\n    rna_data = train_labels[train_labels[\"target_id\"] == target_id]\n    \n    # Plot the RNA structure with edges\n    plot_rna_structure_with_edges(ax, rna_data, target_id)\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T16:38:24.126704Z","iopub.execute_input":"2025-03-14T16:38:24.127055Z","iopub.status.idle":"2025-03-14T16:38:30.904449Z","shell.execute_reply.started":"2025-03-14T16:38:24.127027Z","shell.execute_reply":"2025-03-14T16:38:30.903181Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_labels","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T18:27:42.265489Z","iopub.execute_input":"2025-03-14T18:27:42.26583Z","iopub.status.idle":"2025-03-14T18:27:42.288494Z","shell.execute_reply.started":"2025-03-14T18:27:42.265804Z","shell.execute_reply":"2025-03-14T18:27:42.287595Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\nValidation Labels Sample:\")\nprint(validation_labels.head())\n\nprint(\"\\nTest Sequences Sample:\")\nprint(test_sequences.head())\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-13T21:51:10.088608Z","iopub.execute_input":"2025-03-13T21:51:10.089153Z","iopub.status.idle":"2025-03-13T21:51:10.122351Z","shell.execute_reply.started":"2025-03-13T21:51:10.089106Z","shell.execute_reply":"2025-03-13T21:51:10.121173Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualizing RNA sequence lengths distribution in training data\ntrain_sequences[\"sequence_length\"] = train_sequences[\"sequence\"].apply(len)\n\nplt.figure(figsize=(10, 5))\nsns.histplot(train_sequences[\"sequence_length\"], bins=30, kde=True)\nplt.xlabel(\"Sequence Length\")\nplt.ylabel(\"Frequency\")\nplt.title(\"Distribution of RNA Sequence Lengths in Training Set\")\nplt.show()\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-13T21:51:11.270842Z","iopub.execute_input":"2025-03-13T21:51:11.271221Z","iopub.status.idle":"2025-03-13T21:51:11.625731Z","shell.execute_reply.started":"2025-03-13T21:51:11.271193Z","shell.execute_reply":"2025-03-13T21:51:11.624836Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualizing the 3D structure of a sample RNA (plotting C1' coordinates)\nsample_rna = train_labels[train_labels[\"ID\"].str.contains(\"1RHT_A\")]\n\nprint(sample_rna.head())\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-13T21:51:14.893565Z","iopub.execute_input":"2025-03-13T21:51:14.89392Z","iopub.status.idle":"2025-03-13T21:51:14.949539Z","shell.execute_reply.started":"2025-03-13T21:51:14.893887Z","shell.execute_reply":"2025-03-13T21:51:14.948728Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fig = plt.figure(figsize=(8, 6))\nax = fig.add_subplot(111, projection='3d')\nax.scatter(sample_rna[\"x_1\"], sample_rna[\"y_1\"], sample_rna[\"z_1\"], c='b', marker='o')\n\nax.set_xlabel(\"X Coordinate (Å)\")\nax.set_ylabel(\"Y Coordinate (Å)\")\nax.set_zlabel(\"Z Coordinate (Å)\")\nax.set_title(\"3D Structure of Sample RNA (C1' Atom Coordinates)\")\n\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-13T21:51:15.535445Z","iopub.execute_input":"2025-03-13T21:51:15.535762Z","iopub.status.idle":"2025-03-13T21:51:15.745818Z","shell.execute_reply.started":"2025-03-13T21:51:15.535733Z","shell.execute_reply":"2025-03-13T21:51:15.744841Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom mpl_toolkits.mplot3d import Axes3D\nfrom sklearn.impute import KNNImputer\nfrom collections import Counter\n\n# -------------------------\n# 1️⃣ Sequence & Length Analysis\n# -------------------------\n\n# Compute sequence lengths\ntrain_sequences[\"sequence_length\"] = train_sequences[\"sequence\"].apply(len)\n\n# Plot sequence length distribution\nplt.figure(figsize=(10, 5))\nsns.histplot(train_sequences[\"sequence_length\"], bins=30, kde=True)\nplt.xlabel(\"Sequence Length\")\nplt.ylabel(\"Frequency\")\nplt.title(\"Distribution of RNA Sequence Lengths in Training Set\")\nplt.show()\n\n# GC-content calculation (proportion of G and C nucleotides)\ndef gc_content(seq):\n    return (seq.count(\"G\") + seq.count(\"C\")) / len(seq)\n\ntrain_sequences[\"GC_content\"] = train_sequences[\"sequence\"].apply(gc_content)\n\n# Scatter plot of sequence length vs GC content\nplt.figure(figsize=(10, 5))\nsns.scatterplot(x=train_sequences[\"sequence_length\"], y=train_sequences[\"GC_content\"])\nplt.xlabel(\"Sequence Length\")\nplt.ylabel(\"GC Content\")\nplt.title(\"GC Content vs Sequence Length\")\nplt.show()\n\n# Nucleotide frequency analysis\nnucleotide_counts = Counter(\"\".join(train_sequences[\"sequence\"]))\nplt.figure(figsize=(8, 5))\nsns.barplot(x=list(nucleotide_counts.keys()), y=list(nucleotide_counts.values()))\nplt.xlabel(\"Nucleotide\")\nplt.ylabel(\"Frequency\")\nplt.title(\"Nucleotide Frequency in RNA Sequences\")\nplt.show()\n\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-13T21:51:24.398948Z","iopub.execute_input":"2025-03-13T21:51:24.399266Z","iopub.status.idle":"2025-03-13T21:51:25.4383Z","shell.execute_reply.started":"2025-03-13T21:51:24.399245Z","shell.execute_reply":"2025-03-13T21:51:25.437294Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\n# Load train_sequences again (after reset)\ntrain_sequences = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_sequences.csv\")\n\n# Compute sequence lengths\ntrain_sequences[\"sequence_length\"] = train_sequences[\"sequence\"].apply(len)\n\n# Separate outliers (length > 200)\noutliers = train_sequences[train_sequences[\"sequence_length\"] > 200]\nfiltered_sequences = train_sequences[train_sequences[\"sequence_length\"] <= 200]\n\n# Plot sequence length distribution without outliers\nplt.figure(figsize=(10, 5))\nsns.histplot(filtered_sequences[\"sequence_length\"], bins=30, kde=True)\nplt.xlabel(\"Sequence Length\")\nplt.ylabel(\"Frequency\")\nplt.title(\"Distribution of RNA Sequence Lengths (<= 1000 residues)\")\nplt.show()\n\n# Plot outliers separately\nplt.figure(figsize=(10, 4))\nsns.histplot(outliers[\"sequence_length\"], bins=10, color='red')\nplt.xlabel(\"Sequence Length\")\nplt.ylabel(\"Frequency\")\nplt.title(\"Outliers: RNA Sequences with Length > 1000\")\nplt.show()\n\n# Display the outlier data\noutliers.reset_index(drop=True, inplace=True)\noutliers.head(10)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T20:10:51.624294Z","iopub.execute_input":"2025-04-14T20:10:51.624614Z","iopub.status.idle":"2025-04-14T20:10:52.067449Z","shell.execute_reply.started":"2025-04-14T20:10:51.624588Z","shell.execute_reply":"2025-04-14T20:10:52.066767Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------\n# 3️⃣ Missing Data Handling (Imputation)\n# -------------------------\n\n# Checking missing values\nmissing_values = train_labels.isnull().sum()\nprint(\"Missing Values:\\n\", missing_values)\n\n# Using KNN Imputer to fill missing 3D coordinates\nfrom sklearn.impute import SimpleImputer\n\n# Creating a mean imputer\nmean_imputer = SimpleImputer(strategy=\"mean\")\n\n# Apply mean imputation for missing values\ntrain_labels[[\"x_1\", \"y_1\", \"z_1\"]] = mean_imputer.fit_transform(train_labels[[\"x_1\", \"y_1\", \"z_1\"]])\n\n# Check if missing values are handled\nprint(\"Missing Values After Imputation:\\n\", train_labels.isnull().sum())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-13T21:51:25.439751Z","iopub.execute_input":"2025-03-13T21:51:25.440121Z","iopub.status.idle":"2025-03-13T21:51:25.507569Z","shell.execute_reply.started":"2025-03-13T21:51:25.440088Z","shell.execute_reply":"2025-03-13T21:51:25.506458Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------\n# 4️⃣ Multi-Conformation Handling (Grouping RNA Structures)\n# -------------------------\n\n# Count number of conformations per RNA sequence\nconformation_counts = train_labels[\"ID\"].value_counts()\n\nplt.figure(figsize=(10, 5))\nsns.histplot(conformation_counts, bins=30, kde=True)\nplt.xlabel(\"Number of Conformations\")\nplt.ylabel(\"Frequency\")\nplt.title(\"Distribution of Multiple Conformations Per RNA Sequence\")\nplt.show()\n\n# Print some examples of RNA sequences with multiple conformations\nmulti_conformations = conformation_counts[conformation_counts > 1].index[:5]\nprint(\"Examples of RNA sequences with multiple conformations:\", multi_conformations.tolist())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-13T21:51:26.289757Z","iopub.execute_input":"2025-03-13T21:51:26.290072Z","iopub.status.idle":"2025-03-13T21:51:26.68883Z","shell.execute_reply.started":"2025-03-13T21:51:26.290047Z","shell.execute_reply":"2025-03-13T21:51:26.68806Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import tensorflow as tf\nimport numpy as np\nimport pandas as pd\n\n# -----------------------------\n# Load and Process Data\n# -----------------------------\n\ntrain_labels = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_labels.csv\")\ntrain_labels[\"target_id\"] = train_labels[\"ID\"].apply(lambda x: \"_\".join(x.split(\"_\")[:2]))\ntrain_labels = train_labels.sort_values(by=[\"target_id\", \"resid\"])\n\nMAX_SEQ_LEN = 20\nNUCLEOTIDE_MAP = {\"A\": [1, 0, 0, 0], \"C\": [0, 1, 0, 0], \"G\": [0, 0, 1, 0], \"U\": [0, 0, 0, 1]}\n\n\n# Updated version of the create_sequence_chunks function with sliding window support\ndef create_sequence_chunks_with_sliding(data, max_seq_len=20, stride=5):\n    \"\"\"\n    Split RNA sequences into chunks of `max_seq_len` residues using sliding windows.\n    \n    Parameters:\n    - data: DataFrame with RNA structural info.\n    - max_seq_len: Number of residues per chunk.\n    - stride: Number of residues to move the window each time.\n    \n    Returns:\n    - A list of chunks, each containing a (max_seq_len, 5) matrix [resname, resid, x, y, z].\n    \"\"\"\n    sequences = []\n    grouped = data.groupby(\"target_id\")\n\n    for target_id, group in grouped:\n        coords = group[[\"resname\", \"resid\", \"x_1\", \"y_1\", \"z_1\"]].values\n\n        for i in range(0, len(coords), stride):\n            chunk = coords[i: i + max_seq_len]\n            if len(chunk) < max_seq_len:\n                padding = np.array([[\"N\", -1, np.nan, np.nan, np.nan]] * (max_seq_len - len(chunk)))\n                chunk = np.vstack((chunk, padding))\n            sequences.append(chunk)\n\n    return sequences\n\n\n\n\nclass RNATensorflowDataset(tf.keras.utils.Sequence):\n    def __init__(self, sequences, batch_size=32):\n        self.sequences = sequences\n        self.batch_size = batch_size\n\n    def __len__(self):\n        return int(90)\n\n\n    def __getitem__(self, idx):\n        batch_chunks = self.sequences[idx * self.batch_size:(idx + 1) * self.batch_size]\n\n        node_features_batch = []\n        positions_batch = []\n        mask_indices = []\n        targets = []\n\n        for chunk in batch_chunks:\n            node_features = []\n            positions = []\n\n            # Choose a valid index to mask\n            valid_indices = [i for i, row in enumerate(chunk) if row[0] in NUCLEOTIDE_MAP]\n            masked_idx = np.random.choice(valid_indices) if valid_indices else 0\n            target = chunk[masked_idx, 2:5].astype(np.float32)\n\n            for i, row in enumerate(chunk):\n                nucleotide = row[0]\n                pos_index = i / MAX_SEQ_LEN\n                is_masked = 1.0 if i == masked_idx else 0.0\n                coords = np.array([0.0, 0.0, 0.0], dtype=np.float32) if i == masked_idx else row[2:5].astype(np.float32)\n            \n                one_hot = NUCLEOTIDE_MAP.get(nucleotide, [0, 0, 0, 0])\n                features = one_hot + [pos_index, is_masked] + coords.tolist()\n                node_features.append(features)\n                positions.append(coords)\n\n\n            node_features_batch.append(node_features)\n            positions_batch.append(positions)\n            mask_indices.append(masked_idx)\n            targets.append(target)\n\n        return (\n            tf.convert_to_tensor(node_features_batch, dtype=tf.float32),  # shape (B, 20, 9)\n            tf.convert_to_tensor(positions_batch, dtype=tf.float32),      # shape (B, 20, 3)\n            tf.convert_to_tensor(mask_indices, dtype=tf.int32),           # shape (B,)\n            tf.convert_to_tensor(targets, dtype=tf.float32),              # shape (B, 3)\n        )\n\n\n# Create sequence chunks and dataset\n# Create sliding window sequences using stride=5\nsliding_sequence_chunks = create_sequence_chunks_with_sliding(train_labels, max_seq_len=20, stride=5)\n\n# Check how many samples are now created\nnum_samples = len(sliding_sequence_chunks)\nprint(f\"Total training samples with sliding window: {num_samples}\")\nsliding_sequence_chunks[:1]  # Show the first chunk for inspection\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T20:36:37.7395Z","iopub.execute_input":"2025-04-16T20:36:37.739839Z","iopub.status.idle":"2025-04-16T20:36:38.555393Z","shell.execute_reply.started":"2025-04-16T20:36:37.739813Z","shell.execute_reply":"2025-04-16T20:36:38.554732Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_sequence_chunks(data):\n    \"\"\"Split RNA sequences into chunks of MAX_SEQ_LEN residues.\"\"\"\n    sequences = []\n    grouped = data.groupby(\"target_id\")\n\n    for target_id, group in grouped:\n        coords = group[[\"resname\", \"resid\", \"x_1\", \"y_1\", \"z_1\"]].values\n\n        for i in range(0, len(coords), MAX_SEQ_LEN):\n            chunk = coords[i: i + MAX_SEQ_LEN]\n            if len(chunk) < MAX_SEQ_LEN:\n                padding = np.array([[\"N\", -1, np.nan, np.nan, np.nan]] * (MAX_SEQ_LEN - len(chunk)))\n                chunk = np.vstack((chunk, padding))\n            sequences.append(chunk)\n\n    return sequences\n\n\n# Create sequence chunks and dataset\nsequence_chunks = create_sequence_chunks(train_labels)\nrna_tf_dataset = RNATensorflowDataset(sequence_chunks, batch_size=128)\n\n# Preview a batch\nsample_batch = rna_tf_dataset[0]\n\n\n# Convert tensors to numpy\nnode_features_np = sample_batch[0].numpy()\npositions_np = sample_batch[1].numpy()\nmasked_indices_np = sample_batch[2].numpy()\ntargets_np = sample_batch[3].numpy()\n\n# Flatten for display\nimport pandas as pd\n\nflat_batch = {\n    \"Masked Index\": masked_indices_np.tolist(),\n    \"Target Coordinates\": targets_np.tolist()\n}\n\n# Show as DataFrame\npd.DataFrame(flat_batch)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T20:36:38.556551Z","iopub.execute_input":"2025-04-16T20:36:38.556888Z","iopub.status.idle":"2025-04-16T20:36:39.082675Z","shell.execute_reply.started":"2025-04-16T20:36:38.556855Z","shell.execute_reply":"2025-04-16T20:36:39.081957Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import tensorflow as tf\n\nclass EGNNLayer(tf.keras.layers.Layer):\n    def __init__(self, hidden_dim):\n        super(EGNNLayer, self).__init__()\n        self.hidden_dim = hidden_dim\n\n        # MLP to compute messages based on feature differences and squared distance\n        self.message_mlp = tf.keras.Sequential([\n            tf.keras.layers.Dense(hidden_dim, activation='relu'),\n            tf.keras.layers.Dense(hidden_dim)\n        ])\n\n        # Coordinate update MLP\n        self.coord_mlp = tf.keras.Sequential([\n            tf.keras.layers.Dense(1, activation='relu'),  # One scalar to scale directional vector\n        ])\n\n        # Node feature update MLP\n        self.feature_mlp = tf.keras.Sequential([\n            tf.keras.layers.Dense(hidden_dim, activation='relu'),\n            tf.keras.layers.Dense(hidden_dim)\n        ])\n\n    def call(self, node_features, positions, edge_index=None):\n        \"\"\"\n        Arguments:\n        - node_features: (B, N, F)\n        - positions: (B, N, 3)\n        Returns:\n        - updated node_features: (B, N, F)\n        - updated positions: (B, N, 3)\n        \"\"\"\n        B, N, F = node_features.shape\n\n        # Step 1: Create all pairwise differences (broadcasted)\n        pos_i = tf.expand_dims(positions, 2)  # (B, N, 1, 3)\n        pos_j = tf.expand_dims(positions, 1)  # (B, 1, N, 3)\n        diff = pos_i - pos_j  # (B, N, N, 3)\n        dist2 = tf.reduce_sum(tf.square(diff), axis=-1, keepdims=True)  # (B, N, N, 1)\n\n        # Step 2: Message computation\n        h_i = tf.expand_dims(node_features, 2)  # (B, N, 1, F)\n        h_j = tf.expand_dims(node_features, 1)  # (B, 1, N, F)\n        message_input = tf.concat([h_i - h_j, dist2], axis=-1)  # (B, N, N, F+1)\n\n        messages = self.message_mlp(message_input)  # (B, N, N, F)\n\n        # Step 3: Aggregate messages\n        agg_messages = tf.reduce_sum(messages, axis=2)  # (B, N, F)\n\n        # Step 4: Update node features\n        updated_features = self.feature_mlp(agg_messages)  # (B, N, F)\n\n        # Step 5: Update coordinates\n        coord_weights = self.coord_mlp(messages)  # (B, N, N, 1)\n        coord_update = tf.reduce_sum(coord_weights * diff, axis=2)  # (B, N, 3)\n        updated_positions = positions + coord_update  # (B, N, 3)\n\n        return updated_features, updated_positions\n\nimport tensorflow as tf\n\nclass RNAGNN(tf.keras.Model):\n    def __init__(self, hidden_dim=64, num_layers=3):\n        super(RNAGNN, self).__init__()\n        self.hidden_dim = hidden_dim\n        self.num_layers = num_layers\n\n        # Initial embedding layer for node features (input_dim = 9 → hidden_dim)\n        self.embed = tf.keras.layers.Dense(hidden_dim)\n\n        # Stack multiple EGNN layers\n        self.egnn_layers = [EGNNLayer(hidden_dim) for _ in range(num_layers)]\n\n        # MLP head to predict (x, y, z)\n        self.mlp = tf.keras.Sequential([\n            tf.keras.layers.Dense(128, activation='relu'),\n            tf.keras.layers.Dense(3)  # output 3D coordinates\n        ])\n\n    def call(self, node_features, positions, masked_idx):\n        \"\"\"\n        Arguments:\n        - node_features: (B, N, 9) including one-hot, pos index, mask flag, coords\n        - positions: (B, N, 3)\n        - masked_idx: (B,) index of masked residue per sequence\n\n        Returns:\n        - predicted_coords: (B, 3) predicted (x, y, z) for each masked residue\n        \"\"\"\n        x = self.embed(node_features)  # shape (B, N, hidden_dim)\n\n        for egnn in self.egnn_layers:\n            x, positions = egnn(x, positions)\n\n        # Gather masked embeddings for each sample in the batch\n        # masked_idx: (B,) → convert to (B, 1) → gather along dim=1\n        masked_repr = tf.gather(x, tf.expand_dims(masked_idx, axis=1), batch_dims=1)\n        masked_repr = tf.squeeze(masked_repr, axis=1)  # shape: (B, hidden_dim)\n\n        # Predict (x, y, z)\n        predicted_coords = self.mlp(masked_repr)  # shape: (B, 3)\n\n        return predicted_coords\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T20:36:39.083762Z","iopub.execute_input":"2025-04-16T20:36:39.084019Z","iopub.status.idle":"2025-04-16T20:36:39.094034Z","shell.execute_reply.started":"2025-04-16T20:36:39.083999Z","shell.execute_reply":"2025-04-16T20:36:39.093092Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define and create the RNA GNN model object\nhidden_dim = 64\nnum_layers = 10\n\n# Instantiate the model\nrna_model = RNAGNN(hidden_dim=hidden_dim, num_layers=num_layers)\n\n# Dummy input to build the model and show summary\nB, N, F = 4, 20, 9  # batch size, number of nodes, feature dimension\ndummy_node_features = tf.random.normal((B, N, F))\ndummy_positions = tf.random.normal((B, N, 3))\ndummy_masked_idx = tf.constant([3, 7, 12, 5], dtype=tf.int32)  # one masked node per sequence\n\n# Call the model once to build it\n_ = rna_model(dummy_node_features, dummy_positions, dummy_masked_idx)\n\n# Show the model summary\nrna_model.summary()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T20:36:39.354208Z","iopub.execute_input":"2025-04-16T20:36:39.354526Z","iopub.status.idle":"2025-04-16T20:36:40.779954Z","shell.execute_reply.started":"2025-04-16T20:36:39.35449Z","shell.execute_reply":"2025-04-16T20:36:40.77917Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\n# Shuffle and split the data\ntrain_chunks, val_chunks = train_test_split(sliding_sequence_chunks, test_size=0.1, random_state=42)\n\n# Re-initialize data loaders\ntrain_dataset = RNATensorflowDataset(train_chunks, batch_size=256)\nval_dataset = RNATensorflowDataset(val_chunks, batch_size=256)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T20:36:40.781174Z","iopub.execute_input":"2025-04-16T20:36:40.781489Z","iopub.status.idle":"2025-04-16T20:36:40.794343Z","shell.execute_reply.started":"2025-04-16T20:36:40.781466Z","shell.execute_reply":"2025-04-16T20:36:40.793587Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(len(val_dataset))  # Number of chunks (i.e., batches)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T20:36:46.915101Z","iopub.execute_input":"2025-04-16T20:36:46.915424Z","iopub.status.idle":"2025-04-16T20:36:46.92004Z","shell.execute_reply.started":"2025-04-16T20:36:46.915397Z","shell.execute_reply":"2025-04-16T20:36:46.919113Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\nimport tensorflow as tf\nimport pandas as pd\nimport numpy as np\n\n# Re-initialize model\nrna_model = RNAGNN(hidden_dim=64, num_layers=30)\n\nloss_fn = tf.keras.losses.MeanSquaredError()\noptimizer = tf.keras.optimizers.Adam(learning_rate=1e-3)\n\nepochs = 50\ntrain_logs = []\n\nfor epoch in range(epochs):\n    total_loss = 0.0\n    total_l2_error = 0.0\n    steps = 0\n\n    print(f\"\\nEpoch {epoch + 1}/{epochs}\")\n    \n\n    for batch in tqdm(train_dataset, desc=f\"Training Epoch {epoch+1}\", total=len(train_dataset), leave=False):\n\n        node_features, positions, masked_idx, targets = batch\n        if node_features.shape[0] == 0:\n            continue\n\n        with tf.GradientTape() as tape:\n            predictions = rna_model(node_features, positions, masked_idx)\n            loss = loss_fn(targets, predictions)\n\n        l2_error = tf.reduce_mean(tf.norm(targets - predictions, axis=1))\n        grads = tape.gradient(loss, rna_model.trainable_variables)\n        optimizer.apply_gradients(zip(grads, rna_model.trainable_variables))\n\n        total_loss += loss.numpy()\n        total_l2_error += l2_error.numpy()\n        steps += 1\n\n    epoch_loss = total_loss / steps\n    epoch_l2_error = total_l2_error / steps\n\n    # Validation phase\n    val_loss = 0.0\n    val_l2_error = 0.0\n    val_steps = 0\n    print(\"training completed\")\n    for val_batch in tqdm(val_dataset,desc=f\"Validation Epoch {epoch+1}\", leave=False):\n        print(\"validation loop\")\n        node_features, positions, masked_idx, targets = val_batch\n        # if node_features.shape[0] == 0:\n        #     continue\n\n        val_preds = rna_model(node_features, positions, masked_idx)\n        val_batch_loss = loss_fn(targets, val_preds)\n        val_l2 = tf.reduce_mean(tf.norm(targets - val_preds, axis=1))\n\n        val_loss += val_batch_loss.numpy()\n        val_l2_error += val_l2.numpy()\n        val_steps += 1\n\n    avg_val_loss = val_loss / val_steps\n    avg_val_l2 = val_l2_error / val_steps\n\n    train_logs.append({\n        \"epoch\": epoch + 1,\n        \"train_loss\": epoch_loss,\n        \"train_l2_error\": epoch_l2_error,\n        \"val_loss\": avg_val_loss,\n        \"val_l2_error\": avg_val_l2\n    })\n\n    print(f\"Epoch {epoch+1}/{epochs} - \"\n          f\"Train Loss: {epoch_loss:.4f}, L2: {epoch_l2_error:.4f} | \"\n          f\"Val Loss: {avg_val_loss:.4f}, L2: {avg_val_l2:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T20:32:54.535224Z","iopub.execute_input":"2025-04-16T20:32:54.535532Z","iopub.status.idle":"2025-04-16T20:35:30.690857Z","shell.execute_reply.started":"2025-04-16T20:32:54.535503Z","shell.execute_reply":"2025-04-16T20:35:30.689498Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}