{"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"},{"sourceId":230001578,"sourceType":"kernelVersion"},{"sourceId":230156592,"sourceType":"kernelVersion"},{"sourceId":230156597,"sourceType":"kernelVersion"}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"This notebook explores a naive approach for identifying candidate base pairs that may show covariation across homologous RNA sequences, based on a multiple sequence alignment (MSA), and visualizes whether those pairs are spatially close in the 3D structure.\n\nThe idea is inspired by a discussion in [1], and this notebook simply implements and tries out that concept in a lightweight way.\n\n> 💡 For background on covarying base pairs and their role in RNA structure prediction, please refer to [1].\n\n---\n\n## 💡 Key Insights\n\n1. Base pairs identified using the **pairing score** are often found to be **spatially close** in the 3D structure — supporting their structural relevance.\n2. However, spatial proximity doesn't always match covariation patterns — likely due to offsets between backbone atoms and interaction sites.\n3. This scoring method could potentially help narrow down plausible base-pair candidates, and may be useful when incorporated into RNA folding models.\n\n---\n\n## ✅ References\n\nThis notebook is based on ideas and insights from the following resources:\n\n1. [Discussion: Covarying Base Pairs in RNA Structure Prediction](https://www.kaggle.com/competitions/stanford-rna-3d-folding/discussion/568633) — by @tilii7  \n2. [Minimal EDA Notebook](https://www.kaggle.com/code/tatamikenn/stanford-3d-folding-minimal-eda) — provides preprocessed structural data\n\n---\n\n## 📝 Updates\n\n- v8: original\n- v9: calculate paring score with mutual information (MI)","metadata":{}},{"cell_type":"code","source":"%load_ext autoreload\n%autoreload 2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T03:17:19.463067Z","iopub.execute_input":"2025-03-29T03:17:19.463472Z","iopub.status.idle":"2025-03-29T03:17:19.489084Z","shell.execute_reply.started":"2025-03-29T03:17:19.463441Z","shell.execute_reply":"2025-03-29T03:17:19.488081Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import kagglehub\n\n\nsr3f = kagglehub.package_import('tatamikenn/stanford-rna-3d-folding-utility-packages/versions/3')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T03:17:19.706531Z","iopub.execute_input":"2025-03-29T03:17:19.706945Z","iopub.status.idle":"2025-03-29T03:17:26.540994Z","shell.execute_reply.started":"2025-03-29T03:17:19.706915Z","shell.execute_reply":"2025-03-29T03:17:26.539634Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import polars as pl\n\nmsa = pl.read_parquet(\"/kaggle/input/stanford-3d-folding-minimal-eda/metadata/msa.parquet\")\nmsa = msa.with_columns(\n    pl.col(\"idx\").len().over(\"target_id\").alias(\"num_samples\"),\n    pl.col(\"seq\").str.len_chars().alias(\"seq_len\"),\n)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-03-29T03:17:26.542205Z","iopub.execute_input":"2025-03-29T03:17:26.542507Z","iopub.status.idle":"2025-03-29T03:17:28.864288Z","shell.execute_reply.started":"2025-03-29T03:17:26.542483Z","shell.execute_reply":"2025-03-29T03:17:28.863122Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## How many samples in MSA?","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T03:17:28.865871Z","iopub.execute_input":"2025-03-29T03:17:28.866184Z","iopub.status.idle":"2025-03-29T03:17:28.885520Z","shell.execute_reply.started":"2025-03-29T03:17:28.866158Z","shell.execute_reply":"2025-03-29T03:17:28.884365Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"_, axes = plt.subplots(1, 2, figsize=(13, 4))\nax = axes[0]\nnum_samples = (\n    msa.group_by(\"target_id\")\n    .agg(num_samples=pl.col(\"num_samples\").first())[\"num_samples\"]\n    .to_numpy()\n)\nmean_samples = num_samples.mean()\nax.axvline(mean_samples, color=\"red\", linestyle=\"--\", label=f\"mean: {mean_samples:.0f}\")\nax.hist(num_samples, bins=100, color=\"blue\", alpha=0.7, edgecolor=\"white\")\nax.grid()\nax.set(\n    xlabel=\"Number of samples per target\",\n    ylabel=\"Count\",\n    title=\"Histogram of MSA samples per target\",\n)\nax.legend()\n\nax = axes[1]\n\nnum_samples = (\n    msa.group_by(\"target_id\")\n    .agg(num_samples=pl.col(\"num_samples\").first())\n    .filter(pl.col(\"num_samples\").lt(15))[\"num_samples\"]\n    .to_numpy()\n)\nuniq, count = np.unique(num_samples, return_counts=True)\nax.bar(uniq, count, color=\"blue\", alpha=0.7, edgecolor=\"white\")\nax.grid()\nax.set(\n    xlabel=\"Number of samples per target\",\n    ylabel=\"Count\",\n    title=\"Count of targets with < 15 samples\",\n)\n\nplt.show()","metadata":{"trusted":true,"_kg_hide-input":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-03-28T12:52:41.022094Z","iopub.execute_input":"2025-03-28T12:52:41.022393Z","iopub.status.idle":"2025-03-28T12:52:43.286547Z","shell.execute_reply.started":"2025-03-28T12:52:41.022369Z","shell.execute_reply":"2025-03-28T12:52:43.285400Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.colors as mcolors\nfrom scipy.special import softmax\n\nplot_pdb = sr3f.plot_pdb\nextract_c1_atoms = sr3f.extract_c1_atoms\n\n\ndef calc_seq_chars(target_id):\n    \"\"\"\n    指定した target_id に対して、シーケンスの文字配列を返す。\n    返り値は shape (L, n) となり、各列が1シーケンス、各行が位置を示す。\n    \"\"\"\n    df = msa.filter(pl.col(\"target_id\") == target_id)\n    df = df.with_columns(pl.col(\"seq\").str.split(\"\").alias(\"seqs\"))\n    seq = np.array(df[\"seqs\"].to_list())\n    return seq.T  # 転置して、行: 位置, 列: シーケンス\n\n\ndef calc_pairing_score(seq_chars, normalize=True, zero_diag=True, symmetric=True, temperature=0.1):\n    \"\"\"\n    各位置間のmutual information (MI) をベクトル化により計算してスコアとして返す。\n\n    seq_chars: 文字配列、shape (L, n)\n               各行: 位置, 各列: シーケンス\n               ギャップは '-' として扱い、MI計算の際には除外する。\n    normalize: 各行方向でsoftmaxを計算して正規化する場合 True\n    zero_diag: 自分自身とのペア（対角成分）のスコアを0にする場合 True\n    temperature: softmax計算時の温度パラメータ（デフォルトは1.0）\n\n    出力: shape (L, L) のMIスコア行列\n    \"\"\"\n    import numpy as np\n    from scipy.special import softmax\n\n    # 対象とする塩基。ギャップ('-')は除外する。\n    symbols = np.array([\"A\", \"C\", \"G\", \"U\"])\n    L, n = seq_chars.shape\n\n    # ギャップでないシーケンスのマスク (shape: (L, n))\n    valid_mask = seq_chars != \"-\"\n\n    # one-hot encoding: 各位置・シーケンス・塩基に対して True/False\n    # shape: (L, n, 4)\n    oh = (seq_chars[:, :, None] == symbols[None, None, :]).astype(np.float64)\n\n    # 各位置のマージナルカウント (有効なシンボルのカウント)\n    # shape: (L, 4)\n    marg_counts = oh.sum(axis=1)\n\n    # 各位置の有効なシーケンス数\n    N = valid_mask.sum(axis=1)\n\n    # マージナル確率: p_i(a) = count(a) / (有効シーケンス数)\n    p = np.where(N[:, None] > 0, marg_counts / N[:, None], 0)\n\n    # 各位置ペアで両位置ともギャップでないシーケンス数 (shape: (L, L))\n    valid_counts = valid_mask.astype(np.int64) @ valid_mask.astype(np.int64).T\n\n    # 各位置ペアのjoint counts:\n    # joint_counts[i,j,a,b] = ∑ₖ [oh[i,k,a] * oh[j,k,b]]\n    # shape: (L, L, 4, 4)\n    joint_counts = np.einsum(\"ika,jkb->ijab\", oh, oh)\n\n    # joint probability: p_{ij}(a,b) = joint_counts / (有効シーケンス数)\n    joint_probs = np.where(\n        valid_counts[:, :, None, None] > 0,\n        joint_counts / valid_counts[:, :, None, None],\n        0.0,\n    )\n\n    # 位置ごとの独立分布の積: p_i(a)*p_j(b)\n    prod = p[:, None, :, None] * p[None, :, None, :]\n\n    # MI計算: joint_probs * log(joint_probs / (p_i*p_j)) を全塩基ペアについて足し合わせる\n    with np.errstate(divide=\"ignore\", invalid=\"ignore\"):\n        mi_terms = np.where(\n            joint_probs > 0, joint_probs * np.log(joint_probs / prod), 0.0\n        )\n\n    # 各位置ペア (i,j) のMIは塩基次元での総和\n    mi_matrix = mi_terms.sum(axis=(-2, -1))\n\n    if zero_diag:\n        np.fill_diagonal(mi_matrix, -np.inf)\n\n    if normalize:\n        # scipy.special.softmaxを利用して各行方向にsoftmax正規化 (温度パラメータ付き)\n        mi_matrix = softmax(mi_matrix / temperature, axis=1)\n\n    if symmetric:\n        # 対称化\n        mi_matrix = (mi_matrix + mi_matrix.T) / 2\n\n    return mi_matrix\n\n\ndef plot_pairing_score(sim, threshold=0.5):\n    bin_sim = np.where(sim > threshold, 1, 0)\n    _, axes = plt.subplots(1, 2, figsize=(10, 4))\n    ax = axes[0]\n    ax.imshow(sim, cmap=\"hot\", interpolation=\"nearest\")\n    ax.set(\n        title=\"Pairing Score matrix\",\n        xlabel=\"Index\",\n        ylabel=\"Index\",\n    )\n    plt.colorbar(ax.imshow(sim, cmap=\"hot\", interpolation=\"nearest\"))\n\n    ax = axes[1]\n    ax.imshow(bin_sim, cmap=\"hot\", interpolation=\"nearest\")\n    ax.set(\n        title=f\"Pairing Score > {threshold}\",\n        xlabel=\"Index\",\n        ylabel=\"Index\",\n    )\n    plt.show()\n\n\ndef get_color_lists(color_dict=mcolors.TABLEAU_COLORS):\n    color_list = list(color_dict.values())\n\n    return color_list\n\n\ndef search_pairs(pairing_score_mat, threshold=0.5):\n    paring_indices = np.argmax(pairing_score_mat, axis=1)\n    scores = pairing_score_mat[range(len(pairing_score_mat)), paring_indices]\n    candidate_idxs = np.where(scores > threshold)[0]\n    return candidate_idxs, scores[candidate_idxs], paring_indices[candidate_idxs]\n\n\ndef plot_pairing_scores(candidate_idxs, pairing_score_mat):\n    _, ax = plt.subplots()\n\n    pairs = []\n    for idx in candidate_idxs:\n        pair_idx = np.argmax(pairing_score_mat[:, idx])\n        ax.axvline(\n            pair_idx,\n            color=\"gray\",\n            linestyle=\"--\",\n        )\n        ax.plot(\n            pairing_score_mat[:, idx],\n            label=f\"index={idx}, pair={pair_idx} (score={pairing_score_mat[pair_idx, idx]:.2f})\",\n        )\n        pairs.append((idx, pair_idx))\n\n    ax.legend(bbox_to_anchor=(1.02, 1), loc=\"upper left\")\n    ax.set(\n        title=\"Pairing Score\",\n        xlabel=\"Index\",\n        ylabel=\"Paring Score\",\n    )\n    plt.show()\n\n\ndef visualize_rna(\n    target_id,\n    pairs: list[tuple[int, int]],\n    scores: list[float],\n    width=\"100%\",\n    height=600,\n    data_dir=\"/kaggle/input/stanford-3d-folding-minimal-eda/pdb\",\n):\n    pdb_file = f\"{data_dir}/{target_id}_1.pdb\"\n\n    view = plot_pdb(\n        pdb_file, filter_option=\"C1\", per_index=5, width=width, height=height\n    )\n\n    c1_atoms = extract_c1_atoms(pdb_file)\n\n    # add bonds\n    def add_cylindar(start, end, score):\n        view.addCylinder(\n            {\n                \"start\": {\"x\": start[\"x\"], \"y\": start[\"y\"], \"z\": start[\"z\"]},\n                \"end\": {\"x\": end[\"x\"], \"y\": end[\"y\"], \"z\": end[\"z\"]},\n                \"radius\": 0.1,\n                \"color\": \"cyan\",\n            }\n        )\n        mid = (\n            (start[\"x\"] + end[\"x\"]) / 2,\n            (start[\"y\"] + end[\"y\"]) / 2,\n            (start[\"z\"] + end[\"z\"]) / 2,\n        )\n        view.addLabel(\n            f\"{score:.2f}\",\n            {\n                \"position\": {\"x\": mid[0], \"y\": mid[1], \"z\": mid[2]},\n                \"backgroundColor\": \"black\",\n                \"backgroundOpacity\": 0.3,\n                \"fontColor\": \"white\",\n                \"fontSize\": 14,\n            },\n        )\n\n    for (i, j), score in zip(pairs, scores):\n        add_cylindar(c1_atoms[i], c1_atoms[j], score)\n    view.show()","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2025-03-29T03:39:30.225743Z","iopub.execute_input":"2025-03-29T03:39:30.226182Z","iopub.status.idle":"2025-03-29T03:39:30.266987Z","shell.execute_reply.started":"2025-03-29T03:39:30.226153Z","shell.execute_reply":"2025-03-29T03:39:30.265832Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🛠️ Overview of the Method\n\n### Step 1: Identify candidate base pairs from MSA  \n- Using a multiple sequence alignment (MSA) of homologous RNA sequences, we compute a **pairing score matrix** between columns (i.e., aligned positions).\n- We focus on **anti-diagonals with consistently high pairing scores**, which may indicate potential structural base-pairs.\n\n### Step 2: Visualize 3D atomic positions  \n- We extract the 3D coordinates of the predicted base pairs from the RNA structure.\n- We then visually inspect whether the bases are spatially close, which supports their structural relevance.\n\n## ⚙️ Pairing Score Calculation\n\nLet $ X \\in \\mathbb{C}^{L \\times n} $ be the sequence character matrix, where each row corresponds to an aligned position and each column to a sequence.  \nWe now define the pairing score matrix $ P \\in \\mathbb{R}^{L \\times L} $ based on **mutual information (MI)** between columns, which captures the statistical dependence of base identities across aligned sequences.\n\n### Mutual Information-Based Score\n\n1. **Valid Symbols**:  \n   Only the canonical bases $\\{A, C, G, U\\}$ are considered; gaps (`'-'`) are excluded from the computation.\n\n2. **Mutual Information (MI)**:  \n   The mutual information score between positions $i$ and $j$ is calculated as:\n   $$\n   MI(i, j) = \\sum_{a \\in \\{A,C,G,U\\}} \\sum_{b \\in \\{A,C,G,U\\}} p_{ij}(a, b) \\log \\left( \\frac{p_{ij}(a, b)}{p_i(a) p_j(b)} \\right)\n   $$\n   Contributions from terms with $p_{ij}(a,b) = 0$ are ignored.","metadata":{}},{"cell_type":"markdown","source":"## 1A51_A","metadata":{}},{"cell_type":"code","source":"target_id = \"1A51_A\"\nthreshold = 0.3\nscore_mat = calc_pairing_score(calc_seq_chars(target_id))\nplot_pairing_score(score_mat, threshold=threshold)\ncandidate_idxs, scores, pairing_indices = search_pairs(score_mat, threshold=threshold)\nprint(f\"Pairs with score > {threshold}:\")\nfor i, j, score in zip(candidate_idxs, pairing_indices, scores):\n    print(f\"p({i}, {j})={score:.2f}\")\nplot_pairing_scores(candidate_idxs, score_mat)\nvisualize_rna(target_id, list(zip(candidate_idxs, pairing_indices)), scores)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T03:40:10.436454Z","iopub.execute_input":"2025-03-29T03:40:10.436877Z","iopub.status.idle":"2025-03-29T03:40:12.946952Z","shell.execute_reply.started":"2025-03-29T03:40:10.436846Z","shell.execute_reply":"2025-03-29T03:40:12.945873Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## R1190","metadata":{}},{"cell_type":"code","source":"target_id = \"R1190\"\nthreshold = 0.3\nscore_mat = calc_pairing_score(calc_seq_chars(target_id))\nplot_pairing_score(score_mat, threshold=threshold)\ncandidate_idxs, scores, pairing_indices = search_pairs(score_mat, threshold=threshold)\nprint(f\"Pairs with score > {threshold}:\")\nfor i, j, score in zip(candidate_idxs, pairing_indices, scores):\n    print(f\"p({i}, {j})={score:.2f}\")\n\nvisualize_rna(target_id + \"_A\", list(zip(candidate_idxs, pairing_indices)), scores, height=800)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T03:39:38.593965Z","iopub.execute_input":"2025-03-29T03:39:38.594316Z","iopub.status.idle":"2025-03-29T03:39:42.428078Z","shell.execute_reply.started":"2025-03-29T03:39:38.594292Z","shell.execute_reply":"2025-03-29T03:39:42.426988Z"}},"outputs":[],"execution_count":null}]}