{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.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":118765,"databundleVersionId":15231210,"sourceType":"competition"}],"dockerImageVersionId":31234,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport random\nimport time\nimport warnings\nfrom Bio import pairwise2\nfrom Bio.Seq import Seq\n\nimport torch\nfrom torch import nn\nfrom torch_geometric.nn import GCNConv, GraphNorm\nfrom torch_geometric.data import Data\n\nimport matplotlib.pyplot as plt\nfrom mpl_toolkits.mplot3d import Axes3D\n\nwarnings.filterwarnings('ignore')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T11:53:52.988176Z","iopub.execute_input":"2026-01-31T11:53:52.988496Z","iopub.status.idle":"2026-01-31T11:53:52.996839Z","shell.execute_reply.started":"2026-01-31T11:53:52.988472Z","shell.execute_reply":"2026-01-31T11:53:52.995447Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"NUCLEOTIDE_MAPPING = {\n    'A': 'A', 'U': 'U', 'G': 'G', 'C': 'C',\n    'I': 'A', '1MA': 'A', 'PSU': 'U', 'M2G': 'G', '5MC': 'C', 'T': 'U',\n}\n\ndef clean_sequence(seq):\n    return \"\".join([NUCLEOTIDE_MAPPING.get(b, 'A') for b in seq])\n\nDATA_PATH = '/kaggle/input/stanford-rna-3d-folding-2/'\ntrain_seqs = pd.read_csv(DATA_PATH + 'train_sequences.csv')\ntest_seqs = pd.read_csv(DATA_PATH + 'test_sequences.csv')\ntrain_labels = pd.read_csv(DATA_PATH + 'train_labels.csv')\n\ndef process_labels(labels_df):\n    coords_dict = {}\n    for id_prefix, group in labels_df.groupby(lambda x: labels_df['ID'][x].rsplit('_', 1)[0]):\n        coords_dict[id_prefix] = group.sort_values('resid')[['x_1', 'y_1', 'z_1']].values\n    return coords_dict\n\ntrain_coords_dict = process_labels(train_labels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T11:53:52.998868Z","iopub.execute_input":"2026-01-31T11:53:52.999826Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def adaptive_rna_constraints(coordinates, sequence, confidence=1.0):\n    refined_coords = coordinates.copy()\n    n = len(sequence)\n    strength = 0.68 * (1.0 - min(confidence, 0.96))\n    \n    for _ in range(2): \n        for i in range(n - 1):\n            p1, p2 = refined_coords[i], refined_coords[i+1]\n            dist = np.linalg.norm(p2 - p1)\n            if dist > 0:\n                adj = (5.95 - dist) * strength * 0.45\n                refined_coords[i+1] += (p2 - p1) / dist * adj\n            \n            if i < n - 2:\n                p3 = refined_coords[i+2]\n                dist2 = np.linalg.norm(p3 - p1)\n                if dist2 > 0:\n                    adj2 = (10.2 - dist2) * strength * 0.25\n                    refined_coords[i+2] += (p3 - p1) / dist2 * adj2\n    return refined_coords\n\ndef adapt_template_to_query(query_seq, template_seq, template_coords):\n    q_c = clean_sequence(query_seq)\n    t_c = clean_sequence(template_seq)\n    \n    alignments = pairwise2.align.globalms(Seq(q_c), Seq(t_c), 2, -1, -7, -0.25, one_alignment_only=True)\n    if not alignments: return np.zeros((len(query_seq), 3))\n    \n    a_q, a_t = alignments[0].seqA, alignments[0].seqB\n    new_coords = np.full((len(query_seq), 3), np.nan)\n    q_idx, t_idx = 0, 0\n    \n    for cq, ct in zip(a_q, a_t):\n        if cq != '-' and ct != '-':\n            if t_idx < len(template_coords): new_coords[q_idx] = template_coords[t_idx]\n            q_idx += 1; t_idx += 1\n        elif cq != '-': q_idx += 1\n        elif ct != '-': t_idx += 1\n\n    for i in range(len(new_coords)):\n        if np.isnan(new_coords[i, 0]):\n            prev_v = next((j for j in range(i-1, -1, -1) if not np.isnan(new_coords[j, 0])), -1)\n            next_v = next((j for j in range(i+1, len(new_coords)) if not np.isnan(new_coords[j, 0])), -1)\n            if prev_v >= 0 and next_v >= 0:\n                w = (i - prev_v) / (next_v - prev_v)\n                new_coords[i] = (1-w)*new_coords[prev_v] + w*new_coords[next_v]\n            elif prev_v >= 0: new_coords[i] = new_coords[prev_v] + [3.5, 0, 0]\n            elif next_v >= 0: new_coords[i] = new_coords[next_v] + [3.5, 0, 0]\n            else: new_coords[i] = [i*3.5, 0, 0]\n    return np.nan_to_num(new_coords)\n\n\ndef find_similar_sequences(query_seq, train_seqs_df, train_coords_dict, top_n=5):\n    similar = []\n    q_c = clean_sequence(query_seq)\n    for _, row in train_seqs_df.iterrows():\n        t_id, t_seq = row['target_id'], row['sequence']\n        if t_id not in train_coords_dict: continue\n        if abs(len(t_seq) - len(query_seq)) / max(len(t_seq), len(query_seq)) > 0.4: continue\n        \n        t_c = clean_sequence(t_seq)\n        alns = pairwise2.align.globalms(Seq(q_c), Seq(t_c), 2, -1, -7, -0.25, one_alignment_only=True)\n        if alns:\n            score = alns[0].score / (2 * min(len(query_seq), len(t_seq)))\n            similar.append((t_id, t_seq, score, train_coords_dict[t_id]))\n    \n    similar.sort(key=lambda x: x[2], reverse=True)\n    return similar[:top_n]\n\ndef predict_rna_structures(sequence, target_id, train_seqs_df, train_coords_dict, n_predictions=5):\n    predictions = []\n    similar_seqs = find_similar_sequences(sequence, train_seqs_df, train_coords_dict, top_n=n_predictions)\n    \n    for i in range(n_predictions):\n        if i < len(similar_seqs):\n            t_id, t_seq, sim, t_coords = similar_seqs[i]\n            adapted = adapt_template_to_query(sequence, t_seq, t_coords)\n            refined = adaptive_rna_constraints(adapted, sequence, confidence=sim)\n            \n            noise = 0.0 if i == 0 else max(0.006, (0.38 - sim) * 0.07)\n            if noise > 0: refined += np.random.normal(0, noise, refined.shape)\n            predictions.append(refined)\n        else:\n            n = len(sequence)\n            coords = np.zeros((n, 3))\n            for j in range(1, n): coords[j] = coords[j-1] + [4.0, 0, 0]\n            predictions.append(coords)\n    return predictions","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RNAGNN(nn.Module):\n    def __init__(self, node_dim):\n        super().__init__()\n\n        self.conv1 = GCNConv(node_dim, 128)\n        self.conv2 = GCNConv(128, 128)\n        self.conv3 = GCNConv(128, 64)\n\n        self.norm1 = GraphNorm(128)\n        self.norm2 = GraphNorm(128)\n\n        self.head = nn.Sequential(\n            nn.Linear(64, 64),\n            nn.ReLU(),\n            nn.Linear(64, 3)   # Δx, Δy, Δz\n        )\n\n    def forward(self, x, edge_index):\n        x = self.conv1(x, edge_index)\n        x = self.norm1(x)\n        x = torch.relu(x)\n\n        x = self.conv2(x, edge_index)\n        x = self.norm2(x)\n        x = torch.relu(x)\n\n        x = self.conv3(x, edge_index)\n        return self.head(x)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_graph(sequence, init_coords):\n    n = len(sequence)\n    \n    mapping = {'A':[1,0,0,0], 'U':[0,1,0,0], 'G':[0,0,1,0], 'C':[0,0,0,1]}\n    seq_features = torch.tensor([mapping.get(b, [1,0,0,0]) for b in sequence], dtype=torch.float)\n    \n    pos_tensor = torch.tensor(init_coords, dtype=torch.float)\n    x = torch.cat([seq_features, pos_tensor], dim=-1) \n\n    edges = []\n    for i in range(n-1):\n        edges.append([i, i+1]); edges.append([i+1, i])\n    for i in range(n-2):\n        edges.append([i, i+2]); edges.append([i+2, i])\n\n    edge_index = torch.tensor(edges).t().contiguous()\n    return Data(x=x, edge_index=edge_index, pos=pos_tensor)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def predict_with_gnn(sequence, init_coords, model, device):\n    data = build_graph(sequence, init_coords).to(device)\n    with torch.no_grad():\n        delta = model(data.x, data.edge_index)\n        refined_coords = data.pos + delta\n    return refined_coords.cpu().numpy()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def predict_rna_structures_gnn(sequence, target_id, train_seqs_df, train_coords_dict, model, device, n_predictions=5):\n    predictions = []\n    similar_seqs = find_similar_sequences(sequence, train_seqs_df, train_coords_dict, top_n=n_predictions)\n    \n    for i in range(n_predictions):\n        if i < len(similar_seqs):\n            t_id, t_seq, sim, t_coords = similar_seqs[i]\n            adapted = adapt_template_to_query(sequence, t_seq, t_coords)\n\n            refined = predict_with_gnn(sequence, adapted, model, device)\n            \n            final = adaptive_rna_constraints(refined, sequence, confidence=sim)\n            predictions.append(final)\n        else:\n            n = len(sequence)\n            coords = np.zeros((n, 3))\n            for j in range(1, n): coords[j] = coords[j-1] + [4.0, 0, 0]\n            predictions.append(coords)\n    return predictions","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel = RNAGNN(node_dim=7).to(device) \nmodel.eval()\n\nall_predictions = []\nstart_time = time.time()\n\nfor idx, row in test_seqs.iterrows():\n    if idx % 10 == 0: \n        print(f\"Processing {idx} | {time.time()-start_time:.1f}s\")\n    \n    tid, seq = row['target_id'], row['sequence']\n    \n    preds = predict_rna_structures_gnn(seq, tid, train_seqs, train_coords_dict, model, device)\n    \n    for j in range(len(seq)):\n        res = {'ID': f\"{tid}_{j+1}\", 'resname': seq[j], 'resid': j+1}\n        for i in range(5):\n            res[f'x_{i+1}'], res[f'y_{i+1}'], res[f'z_{i+1}'] = preds[i][j]\n        all_predictions.append(res)\n\nsub = pd.DataFrame(all_predictions)\ncols = ['ID', 'resname', 'resid'] + [f'{c}_{i}' for i in range(1,6) for c in ['x','y','z']]\nsub[cols].to_csv('submission.csv', index=False)\nprint(\"GNN-Enhanced Submission Generated!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_rna_3d(coords, sequence=None, title=\"RNA 3D Structure\"):\n\n    coords = np.asarray(coords)\n    n = coords.shape[0]\n\n    fig = plt.figure(figsize=(8, 7))\n    ax = fig.add_subplot(111, projection='3d')\n\n    colors = np.linspace(0, 1, n)\n\n    sc = ax.scatter(\n        coords[:, 0], coords[:, 1], coords[:, 2],\n        c=colors, cmap='viridis', s=25\n    )\n\n    ax.plot(coords[:, 0], coords[:, 1], coords[:, 2], alpha=0.6)\n\n    ax.set_title(title)\n    ax.set_xlabel(\"X\")\n    ax.set_ylabel(\"Y\")\n    ax.set_zlabel(\"Z\")\n\n    max_range = (coords.max(axis=0) - coords.min(axis=0)).max() / 2\n    mid = coords.mean(axis=0)\n    ax.set_xlim(mid[0] - max_range, mid[0] + max_range)\n    ax.set_ylim(mid[1] - max_range, mid[1] + max_range)\n    ax.set_zlim(mid[2] - max_range, mid[2] + max_range)\n\n    cbar = plt.colorbar(sc, ax=ax, shrink=0.6)\n    cbar.set_label(\"Residue index\")\n\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_rna_3d(preds[0], sequence=seq, title=f\"{tid} | Prediction 1\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_rna_ensemble(predictions, title=\"RNA Ensemble (5 predictions)\"):\n\n    fig = plt.figure(figsize=(8, 7))\n    ax = fig.add_subplot(111, projection='3d')\n\n    cmap = plt.get_cmap(\"tab10\")\n\n    for i, coords in enumerate(predictions):\n        coords = np.asarray(coords)\n        ax.plot(\n            coords[:, 0], coords[:, 1], coords[:, 2],\n            color=cmap(i), alpha=0.7, label=f\"Pred {i+1}\"\n        )\n        ax.scatter(\n            coords[:, 0], coords[:, 1], coords[:, 2],\n            color=cmap(i), s=10\n        )\n\n    ax.set_title(title)\n    ax.set_xlabel(\"X\")\n    ax.set_ylabel(\"Y\")\n    ax.set_zlabel(\"Z\")\n\n    all_coords = np.concatenate(predictions, axis=0)\n    max_range = (all_coords.max(axis=0) - all_coords.min(axis=0)).max() / 2\n    mid = all_coords.mean(axis=0)\n\n    ax.set_xlim(mid[0] - max_range, mid[0] + max_range)\n    ax.set_ylim(mid[1] - max_range, mid[1] + max_range)\n    ax.set_zlim(mid[2] - max_range, mid[2] + max_range)\n\n    ax.legend()\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_rna_ensemble(preds, title=f\"{tid} | GNN Ensemble\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}