{"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":11553390,"sourceType":"competition"}],"dockerImageVersionId":30918,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# CLIP approach\n\nThis notebook is inspired by this paper, where the authors tried a CLIP approach to predict structure of antiobody sequences. Here we try to apply the same to RNA folding. \n\n- https://www.mlsb.io/papers_2023/Enhancing_Antibody_Language_Models_with_Structural_Information.pdf\n","metadata":{}},{"cell_type":"markdown","source":"### Read data","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\n\nfasta_files = []\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        if not\"fasta\" in filename:\n            print(os.path.join(dirname, filename))\n        else:\n            fasta_files.append(filename)\nprint(f\"{len(fasta_files)} fasta files\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:34:14.074072Z","iopub.execute_input":"2025-05-08T12:34:14.074515Z","iopub.status.idle":"2025-05-08T12:34:17.038848Z","shell.execute_reply.started":"2025-05-08T12:34:14.074483Z","shell.execute_reply":"2025-05-08T12:34:17.036934Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sequences_df = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_sequences.csv\")\nsequences_df[[\"target_id\", \"sequence\"]]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:34:17.039915Z","iopub.execute_input":"2025-05-08T12:34:17.040731Z","iopub.status.idle":"2025-05-08T12:34:17.150738Z","shell.execute_reply.started":"2025-05-08T12:34:17.040690Z","shell.execute_reply":"2025-05-08T12:34:17.149276Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"labels_df = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_labels.csv\")\nlabels_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:34:17.152294Z","iopub.execute_input":"2025-05-08T12:34:17.152858Z","iopub.status.idle":"2025-05-08T12:34:17.491533Z","shell.execute_reply.started":"2025-05-08T12:34:17.152813Z","shell.execute_reply":"2025-05-08T12:34:17.490135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/sample_submission.csv\")\nsubmission.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:34:17.492806Z","iopub.execute_input":"2025-05-08T12:34:17.493235Z","iopub.status.idle":"2025-05-08T12:34:17.529572Z","shell.execute_reply.started":"2025-05-08T12:34:17.493201Z","shell.execute_reply":"2025-05-08T12:34:17.528090Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"## Using CLIP approach","metadata":{}},{"cell_type":"markdown","source":"### Read PDB coordinates and store it into a dataframe","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\n\n# Load the CSV (you've already done this)\n# labels_df = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_labels.csv\")\n\n# Drop rows with any missing coordinates\nclean_df = labels_df.dropna(subset=[\"x_1\", \"y_1\", \"z_1\"])\n\n# Group by RNA ID\ngrouped = clean_df.groupby(\"ID\")\n\n# Determine max number of residues across all RNAs\nmax_len = grouped.size().max()\n\n# Function to extract and flatten coordinates for each RNA\ndef extract_flat_coords(group, max_len):\n    coords = group[[\"x_1\", \"y_1\", \"z_1\"]].values  # shape (L, 3)\n    flat = coords.flatten()  # shape (L*3,)\n    # Pad with zeros if sequence is shorter than max_len\n    expected_len = max_len * 3\n    if len(flat) < expected_len:\n        pad = np.zeros(expected_len - len(flat), dtype=np.float32)\n        flat = np.concatenate([flat, pad])\n    return flat.astype(np.float32)\n\n# Apply to all RNA structures\nstruct_features_array = np.stack([\n    extract_flat_coords(group, max_len)\n    for _, group in grouped\n])\n\n# Store the RNA IDs in order (useful for joining later)\nrna_ids = list(grouped.groups.keys())\n\nprint(\"Shape of struct_features_array:\", struct_features_array.shape)\n# Should be (num_sequences, max_len * 3)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:55:05.062423Z","iopub.execute_input":"2025-05-08T12:55:05.062781Z","iopub.status.idle":"2025-05-08T12:56:03.935835Z","shell.execute_reply.started":"2025-05-08T12:55:05.062753Z","shell.execute_reply":"2025-05-08T12:56:03.934407Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Read Sequences, one-hot encode them, and store them into a df","metadata":{}},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import DataLoader, Dataset\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\n\n# Example dataset class: replace with your actual data loading/preprocessing\nclass RNADataset(Dataset):\n    def __init__(self, seq_features, struct_features):\n        self.seq_features = seq_features  # numpy array or list, shape: (N, seq_input_dim)\n        self.struct_features = struct_features  # shape: (N, struct_input_dim)\n    \n    def __len__(self):\n        return len(self.seq_features)\n    \n    def __getitem__(self, idx):\n        return {\n            'seq_features': torch.tensor(self.seq_features[idx], dtype=torch.float),\n            'struct_features': torch.tensor(self.struct_features[idx], dtype=torch.float)\n        }\n\ndef one_hot_encode(seq, max_len=100):\n    mapping = {'A': [1, 0, 0, 0],\n               'C': [0, 1, 0, 0],\n               'G': [0, 0, 1, 0],\n               'U': [0, 0, 0, 1]}\n    encoded = [mapping.get(nt, [0, 0, 0, 0]) for nt in seq.upper()]\n    encoded = encoded[:max_len]  # truncate\n    # pad to max_len\n    encoded += [[0, 0, 0, 0]] * (max_len - len(encoded))\n    return encoded\n\n\n# Create dataset and dataloader\n#dataset = RNADataset(sequences_df[\"sequence\"], labels_df)\n#dataloader = DataLoader(dataset, batch_size=32, shuffle=True)\n# Infer input dimensions from the dataset\n#sample = dataset[0]\n#seq_input_dim = sample['seq_features'].shape[0]\n#struct_input_dim = sample['struct_features'].shape[0]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:51:34.962333Z","iopub.execute_input":"2025-05-08T12:51:34.962728Z","iopub.status.idle":"2025-05-08T12:51:34.970801Z","shell.execute_reply.started":"2025-05-08T12:51:34.962700Z","shell.execute_reply":"2025-05-08T12:51:34.969514Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"seq_input_dim = len(sequences_df)\nstruct_input_dim = len(labels_df)\nprint(seq_input_dim,struct_input_dim )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:51:37.851721Z","iopub.execute_input":"2025-05-08T12:51:37.852050Z","iopub.status.idle":"2025-05-08T12:51:37.858953Z","shell.execute_reply.started":"2025-05-08T12:51:37.852025Z","shell.execute_reply":"2025-05-08T12:51:37.857698Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:51:39.414044Z","iopub.execute_input":"2025-05-08T12:51:39.414470Z","iopub.status.idle":"2025-05-08T12:51:39.420588Z","shell.execute_reply.started":"2025-05-08T12:51:39.414435Z","shell.execute_reply":"2025-05-08T12:51:39.419202Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sequences_encoded = [one_hot_encode(seq) for seq in sequences_df[\"sequence\"]]\nsequences_array = np.array(sequences_encoded, dtype=np.float32)  # shape (N, max_len, 4)\nsequences_flat = sequences_array.reshape(sequences_array.shape[0], -1)  # (N, max_len * 4)\n#sequences_encoded\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:52:38.839875Z","iopub.execute_input":"2025-05-08T12:52:38.840236Z","iopub.status.idle":"2025-05-08T12:52:38.916196Z","shell.execute_reply.started":"2025-05-08T12:52:38.840208Z","shell.execute_reply":"2025-05-08T12:52:38.915149Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Define Model and Data classes","metadata":{}},{"cell_type":"code","source":"\n# Improved version of RNA sequence-structure contrastive learning model\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader\n\n# --------- Encoders --------- #\nclass RNASequenceEncoder(nn.Module):\n    def __init__(self, input_dim, emb_dim):\n        super().__init__()\n        self.encoder = nn.Sequential(\n            nn.Linear(input_dim, 512),\n            nn.ReLU(),\n            nn.Dropout(0.1),\n            nn.Linear(512, emb_dim)\n        )\n    def forward(self, x):\n        return self.encoder(x)\n\nclass RNAStructureEncoder(nn.Module):\n    def __init__(self, input_dim, emb_dim):\n        super().__init__()\n        self.encoder = nn.Sequential(\n            nn.Linear(input_dim, 512),\n            nn.ReLU(),\n            nn.Dropout(0.1),\n            nn.Linear(512, emb_dim)\n        )\n    def forward(self, x):\n        return self.encoder(x)\n\n# --------- Contrastive Loss --------- #\nclass ContrastiveCLIPLoss(nn.Module):\n    def __init__(self, temperature=0.07):\n        super().__init__()\n        self.temperature = temperature\n    def forward(self, seq_embeddings, struct_embeddings):\n        seq_norm = F.normalize(seq_embeddings, dim=1, eps=1e-6)\n        struct_norm = F.normalize(struct_embeddings, dim=1, eps=1e-6)\n        logits = torch.matmul(seq_norm, struct_norm.t()) / self.temperature\n        batch_size = logits.size(0)\n        labels = torch.arange(batch_size).to(logits.device)\n        loss_seq = F.cross_entropy(logits, labels)\n        loss_struct = F.cross_entropy(logits.t(), labels)\n        return (loss_seq + loss_struct) / 2.0\n\n# --------- Training Setup --------- #\nemb_dim = 256\n\nseq_encoder = RNASequenceEncoder(seq_input_dim, emb_dim)\nstruct_encoder = RNAStructureEncoder(struct_input_dim, emb_dim)\ncriterion = ContrastiveCLIPLoss(temperature=0.07)\n\noptimizer = torch.optim.Adam(list(seq_encoder.parameters()) + list(struct_encoder.parameters()), lr=1e-3)\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nsequences_encoded = [one_hot_encode(s) for s in sequences_df[\"sequence\"]]\n\n#dataset = RNADataset(sequences_encoded, struct_features_array)\n\nseq_encoder.to(device)\nstruct_encoder.to(device)\ncriterion.to(device)\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:51:48.941925Z","iopub.execute_input":"2025-05-08T12:51:48.942284Z","iopub.status.idle":"2025-05-08T12:51:49.689139Z","shell.execute_reply.started":"2025-05-08T12:51:48.942253Z","shell.execute_reply":"2025-05-08T12:51:49.687537Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Train","metadata":{}},{"cell_type":"code","source":"# --------- Training Loop --------- #\nnum_epochs = 10\nbest_loss = float('inf')\nfor epoch in range(num_epochs):\n    seq_encoder.train()\n    struct_encoder.train()\n    running_loss = 0.0\n\n    for batch in dataloader:  # Assume DataLoader yields dicts with 'seq_features' and 'struct_features'\n        seq_batch = batch['seq_features'].to(device)\n        struct_batch = batch['struct_features'].to(device)\n\n        optimizer.zero_grad()\n        seq_embeddings = seq_encoder(seq_batch)\n        struct_embeddings = struct_encoder(struct_batch)\n        loss = criterion(seq_embeddings, struct_embeddings)\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item()\n\n    avg_loss = running_loss / len(dataloader)\n    print(f\"Epoch {epoch+1}/{num_epochs}, Loss: {avg_loss:.4f}\")\n\n    # Save best model checkpoint\n    if avg_loss < best_loss:\n        best_loss = avg_loss\n        torch.save({\n            'seq_encoder': seq_encoder.state_dict(),\n            'struct_encoder': struct_encoder.state_dict()\n        }, 'best_clip_model.pt')\n\nprint(\"Training complete. Best loss:\", best_loss)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:40:02.736480Z","iopub.execute_input":"2025-05-08T12:40:02.736810Z","iopub.status.idle":"2025-05-08T12:40:03.510819Z","shell.execute_reply.started":"2025-05-08T12:40:02.736786Z","shell.execute_reply":"2025-05-08T12:40:03.509515Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Predict","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\n\n# Read the test sequences file. Adjust engine/quoting if necessary.\ntest_df = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/test_sequences.csv\", engine=\"python\")\n\n# Prepare a list to collect submission rows\nsubmission_rows = []\n\n# For each test sequence, create one row per residue\nfor _, row in test_df.iterrows():\n    target_id = row[\"target_id\"]\n    sequence = str(row[\"sequence\"]).strip()  # Ensure it's a string and remove extra whitespace/newlines\n    for i, nucleotide in enumerate(sequence):\n        resid = i + 1\n        # Create an ID by appending the residue index to the target_id, e.g., \"R1107_1\"\n        new_id = f\"{target_id}_{resid}\"\n        # For a random submission, fill coordinates with zeros.\n        coords = [0.0] * (3 * 5)  # 5 predictions, each with x, y, z (total 15 numbers)\n        submission_rows.append([new_id, nucleotide, resid] + coords)\n\n# Define column names: ID, resname, resid, followed by x_1, y_1, z_1, ..., x_5, y_5, z_5.\ncolumns = [\"ID\", \"resname\", \"resid\"] + [f\"{axis}_{i}\" for i in range(1, 6) for axis in [\"x\", \"y\", \"z\"]]\n\n# Create the submission DataFrame\nsubmission_df = pd.DataFrame(submission_rows, columns=columns)\n\n# Save the submission file to /kaggle/working (this is the working directory in Kaggle notebooks)\nsubmission_df.to_csv(\"/kaggle/working/submission.csv\", index=False)\nprint(\"Submission file saved to /kaggle/working/submission.csv\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:34:22.192349Z","iopub.status.idle":"2025-05-08T12:34:22.192686Z","shell.execute_reply":"2025-05-08T12:34:22.192554Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T12:34:22.193321Z","iopub.status.idle":"2025-05-08T12:34:22.193631Z","shell.execute_reply":"2025-05-08T12:34:22.193502Z"}},"outputs":[],"execution_count":null}]}