{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":[{"sourceType":"competition","sourceId":118765,"databundleVersionId":15231210},{"sourceType":"datasetVersion","sourceId":10880419,"datasetId":6760509,"databundleVersionId":11247150},{"sourceType":"datasetVersion","sourceId":10880374,"datasetId":6760482,"databundleVersionId":11247092},{"sourceType":"datasetVersion","sourceId":11899194,"datasetId":7479946,"databundleVersionId":12404228},{"sourceType":"datasetVersion","sourceId":11451236,"datasetId":7174725,"databundleVersionId":11892125},{"sourceType":"datasetVersion","sourceId":11230242,"datasetId":7014687,"databundleVersionId":11640305},{"sourceType":"datasetVersion","sourceId":13282339,"datasetId":7162026,"databundleVersionId":13982667}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Stanford RNA 3D Folding – Part 2\n## Fixed + Properly-Improved Pipeline\n\n### Root-cause fixes applied:\n| # | Bug | Fix |\n|---|-----|-----|\n| 1 | Alignment cache returned exhausted iterator | Cache stores **scores only**; alignment re-run fresh when needed |\n| 2 | All 5 preds from 1 template killed diversity | Restored 1-pred-per-template strategy + diversity within each |\n| 3 | `build_segments_map` return-type mismatch | Restored original `(seg_map, stoich_map)` tuple return |\n| 4 | Multi-seed × n_sample = 3× GPU time | Single smart seed strategy; n_sample=5 in one call |\n| 5 | Combined ranking reordered templates wrong | Sort by `(norm_score, pct_id)` tuple — original logic restored |\n| 6 | N_cycle=20 / N_step=400 causes timeout | Kept original 12/250; added optional boost only for short seqs |\n| 7 | Protenix triggered for mid-sim targets | Protenix only when templates genuinely insufficient |\n\n### Genuine improvements kept (bug-free ones only):\n- Direction-aware gap filling in `adapt_template_to_query`\n- Confidence-adaptive refinement passes (+1 pass for conf < 0.3)\n- Soft-diversity within each template slot (noise scale based on sim)\n- Submission format validation cell\n","metadata":{}},{"cell_type":"code","source":"import os\nimport sys\nimport time\nimport json\nimport random\nimport warnings\nimport contextlib\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nwarnings.filterwarnings('ignore')\n\nIS_KAGGLE         = True\nDATA_PATH         = '/kaggle/input/stanford-rna-3d-folding-2/'\nOUTPUT_PATH       = '/kaggle/working/output'\nUSALIGN_BIN       = '/kaggle/working/USalign'\nPROTENIX_DIR      = '/kaggle/working/Protenix'\n\n# ── Config (restored to safe defaults) ───────────────────────\nUSE_PROTENIX           = True\nMIN_SIMILARITY         = 0.15\nMIN_PERCENT_IDENTITY   = 35\nSHOW_VALIDATION        = False\nSHOW_ALIGNMENT_DETAILS = True\nMAKE_SUBMISSION        = True\nDEBUG                  = False\n\nif DEBUG:\n    MIN_SIMILARITY = 0.3\n\ndef seed_everything(seed: int = 42):\n    random.seed(seed)\n    np.random.seed(seed)\n\nseed_everything(42)\n\n! cp /kaggle/input/protenix-packages/packages/USalign /kaggle/working/ 2>/dev/null || true\n! chmod +x /kaggle/working/USalign 2>/dev/null || true\nsys.path.insert(0, '/kaggle/input/rna-3d-utils/')\n\nprint(f\"Data path: {DATA_PATH}\")\nprint(f\"Protenix dir: {PROTENIX_DIR}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T12:13:12.175217Z","iopub.execute_input":"2026-02-19T12:13:12.175541Z","iopub.status.idle":"2026-02-19T12:13:12.454873Z","shell.execute_reply.started":"2026-02-19T12:13:12.175515Z","shell.execute_reply":"2026-02-19T12:13:12.453719Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if IS_KAGGLE:\n    !cp -r /kaggle/input/protenix-packages/packages /kaggle/working\n    %cd /kaggle/working/packages\n    !pip install --no-deps --exists-action=i *.whl\n    %cd /kaggle/working\n\n    !mv /kaggle/working/packages/ihm-2.3/ihm-2.3 /kaggle/working\n    !mv /kaggle/working/packages/modelcif-0.7/modelcif-0.7 /kaggle/working\n    !pip install /kaggle/working/ihm-2.3\n    !pip install /kaggle/working/modelcif-0.7\n    !rm -rf /kaggle/working/ihm-2.3 /kaggle/working/modelcif-0.7\n\n    !pip install /kaggle/input/biopython/biopython-1.85-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n    !pip install /kaggle/input/ml-collections/ml_collections-1.0.0-py3-none-any.whl\n\n    !rm -rf /kaggle/working/packages\n\n    !cp -r /kaggle/input/protenix-mg-packages/protenix_mg_packages /kaggle/working\n    %cd /kaggle/working/protenix_mg_packages\n    !pip install --no-deps --exists-action=i *.whl\n    %cd /kaggle/working\n    !rm -rf /kaggle/working/protenix_mg_packages\n\n    !cp -R /kaggle/input/protenix-rmsa-repo/protenix_kaggle /kaggle/working/\n    !mv protenix_kaggle Protenix\n\n    print(\"All packages installed.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T12:13:12.456478Z","iopub.execute_input":"2026-02-19T12:13:12.456823Z","iopub.status.idle":"2026-02-19T12:13:40.49648Z","shell.execute_reply.started":"2026-02-19T12:13:12.456788Z","shell.execute_reply":"2026-02-19T12:13:40.495581Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json\nimport numpy as np\nimport pandas as pd\nimport time\nimport random\nimport warnings\nimport contextlib\nfrom pathlib import Path\nfrom Bio.Align import PairwiseAligner\n\nwarnings.filterwarnings('ignore')\nseed_everything(42)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T12:13:40.498229Z","iopub.execute_input":"2026-02-19T12:13:40.498466Z","iopub.status.idle":"2026-02-19T12:13:40.503041Z","shell.execute_reply.started":"2026-02-19T12:13:40.498446Z","shell.execute_reply":"2026-02-19T12:13:40.502387Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Loading sequence data...\")\ntrain_seqs        = pd.read_csv(DATA_PATH + 'train_sequences.csv')\nvalidation_seqs   = pd.read_csv(DATA_PATH + 'validation_sequences.csv')\ntest_seqs         = pd.read_csv(DATA_PATH + 'test_sequences.csv')\ntrain_labels      = pd.read_csv(DATA_PATH + 'train_labels.csv')\nvalidation_labels = pd.read_csv(DATA_PATH + 'validation_labels.csv')\n\nprint(f\"Loaded {len(train_seqs)} training sequences\")\nprint(f\"Loaded {len(validation_seqs)} validation sequences\")\nprint(f\"Loaded {len(test_seqs)} test sequences\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T12:13:40.504296Z","iopub.execute_input":"2026-02-19T12:13:40.504557Z","iopub.status.idle":"2026-02-19T12:13:48.795094Z","shell.execute_reply.started":"2026-02-19T12:13:40.504536Z","shell.execute_reply":"2026-02-19T12:13:48.794201Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ══════════════════════════════════════════════════════════════\n# ALIGNMENT ENGINE\n# FIX #1: Score cache stores FLOAT only, NOT the alignment object.\n#          adapt_template_to_query() always runs a fresh alignment.\n#          This prevents exhausted-iterator coordinates collapsing to 0.\n# ══════════════════════════════════════════════════════════════\n\ndef make_aligner():\n    al = PairwiseAligner()\n    al.mode = 'global'\n    al.match_score = 3.0\n    al.mismatch_score = -2.0\n    al.open_gap_score = -10\n    al.extend_gap_score = -0.5\n    al.query_left_open_gap_score   = -5\n    al.query_left_extend_gap_score = -0.5\n    al.query_right_open_gap_score  = -5\n    al.query_right_extend_gap_score = -0.5\n    al.target_left_open_gap_score   = -5\n    al.target_left_extend_gap_score = -0.5\n    al.target_right_open_gap_score  = -5\n    al.target_right_extend_gap_score = -0.5\n    return al\n\n_aligner = make_aligner()\n\n# Score-only cache: fast O(1) lookup for repeated pairs,\n# but NEVER caches the alignment object itself (which is a one-shot iterator).\n_score_cache = {}\n\ndef cached_score(seq1: str, seq2: str) -> float:\n    key = (seq1, seq2)\n    if key not in _score_cache:\n        _score_cache[key] = _aligner.score(seq1, seq2)\n    return _score_cache[key]\n\ndef fresh_alignment(seq1: str, seq2: str):\n    \"\"\"Always returns a fresh alignment object. Never cached — safe to iterate.\"\"\"\n    return next(iter(_aligner.align(seq1, seq2)))\n\n\n# ══════════════════════════════════════════════════════════════\n# FASTA / STOICHIOMETRY / SEGMENT HELPERS\n# FIX #3: build_segments_map returns (seg_map, stoich_map) tuple\n#          to match original signature used everywhere downstream.\n# ══════════════════════════════════════════════════════════════\n\ndef parse_fasta(fasta_content: str):\n    out, cur, seq_parts = {}, None, []\n    for line in str(fasta_content).splitlines():\n        line = line.strip()\n        if not line: continue\n        if line.startswith('>'):\n            if cur is not None:\n                out[cur] = ''.join(seq_parts)\n            cur = line[1:].split()[0]\n            seq_parts = []\n        else:\n            seq_parts.append(line.replace(' ', ''))\n    if cur is not None:\n        out[cur] = ''.join(seq_parts)\n    return out\n\ndef parse_stoichiometry(stoich: str):\n    if pd.isna(stoich) or str(stoich).strip() == '':\n        return []\n    out = []\n    for part in str(stoich).split(';'):\n        ch, cnt = part.split(':')\n        out.append((ch.strip(), int(cnt)))\n    return out\n\ndef get_chain_segments(row):\n    seq    = row['sequence']\n    stoich = row.get('stoichiometry', '')\n    all_seq = row.get('all_sequences', '')\n    if pd.isna(stoich) or pd.isna(all_seq) or str(stoich).strip() == '':\n        return [(0, len(seq))]\n    try:\n        chain_dict = parse_fasta(all_seq)\n        order = parse_stoichiometry(stoich)\n        segs, pos = [], 0\n        for ch, cnt in order:\n            base = chain_dict.get(ch)\n            if base is None: return [(0, len(seq))]\n            for _ in range(cnt):\n                segs.append((pos, pos + len(base)))\n                pos += len(base)\n        return segs if pos == len(seq) else [(0, len(seq))]\n    except Exception:\n        return [(0, len(seq))]\n\ndef build_segments_map(df):\n    \"\"\"Returns (seg_map, stoich_map) — original signature preserved.\"\"\"\n    seg_map, stoich_map = {}, {}\n    for _, r in df.iterrows():\n        tid = r['target_id']\n        seg_map[tid]    = get_chain_segments(r)\n        stoich_map[tid] = str(r.get('stoichiometry', '')\n                              if not pd.isna(r.get('stoichiometry', '')) else '')\n    return seg_map, stoich_map\n\ndef process_labels(labels_df):\n    coords_dict = {}\n    prefixes = labels_df['ID'].str.rsplit('_', n=1).str[0]\n    for id_prefix, group in labels_df.groupby(prefixes):\n        coords_dict[id_prefix] = group.sort_values('resid')[['x_1','y_1','z_1']].values\n    return coords_dict\n\nprint(\"Utility functions ready.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T12:13:48.795993Z","iopub.execute_input":"2026-02-19T12:13:48.796248Z","iopub.status.idle":"2026-02-19T12:13:48.813673Z","shell.execute_reply.started":"2026-02-19T12:13:48.796214Z","shell.execute_reply":"2026-02-19T12:13:48.812964Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Processing coordinates...\")\ntrain_coords_dict    = process_labels(train_labels)\n\ncombined_seqs        = pd.concat([train_seqs, validation_seqs], ignore_index=True)\ncombined_labels      = pd.concat([train_labels, validation_labels], ignore_index=True)\ncombined_coords_dict = process_labels(combined_labels)\n\nprint(f\"Processed {len(train_coords_dict)} training structures\")\nprint(f\"Processed {len(combined_coords_dict)} combined (train+val) structures\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T12:13:48.814544Z","iopub.execute_input":"2026-02-19T12:13:48.814846Z","iopub.status.idle":"2026-02-19T12:14:40.28217Z","shell.execute_reply.started":"2026-02-19T12:13:48.814815Z","shell.execute_reply":"2026-02-19T12:14:40.281376Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"validation_segments_map, _ = build_segments_map(validation_seqs)\ntest_segments_map, _        = build_segments_map(test_seqs)\n\nprint(f\"Built segment maps: {len(validation_segments_map)} validation, {len(test_segments_map)} test\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T12:14:40.28305Z","iopub.execute_input":"2026-02-19T12:14:40.283281Z","iopub.status.idle":"2026-02-19T12:14:40.292598Z","shell.execute_reply.started":"2026-02-19T12:14:40.283261Z","shell.execute_reply":"2026-02-19T12:14:40.291871Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ══════════════════════════════════════════════════════════════\n# SEQUENCE SEARCH\n# FIX #5: Sort key restored to (norm_score, pct_id) — original logic.\n#          Combined ranking metric caused wrong template selection for\n#          moderate-sim / low-identity distant homologs.\n#\n# IMPROVEMENT KEPT: cached_score() for the pre-screening pass\n#   → fast O(1) skip for pairs already scored below threshold\n# ══════════════════════════════════════════════════════════════\n\ndef _build_aligned_strings(query_seq, template_seq, alignment):\n    q_segs, t_segs = alignment.aligned\n    aq, at, qi, ti = [], [], 0, 0\n    for (qs, qe), (ts, te) in zip(q_segs, t_segs):\n        while qi < qs: aq.append(query_seq[qi]); at.append('-'); qi += 1\n        while ti < ts: aq.append('-'); at.append(template_seq[ti]); ti += 1\n        for qp, tp in zip(range(qs, qe), range(ts, te)):\n            aq.append(query_seq[qp]); at.append(template_seq[tp])\n        qi, ti = qe, te\n    while qi < len(query_seq): aq.append(query_seq[qi]); at.append('-'); qi += 1\n    while ti < len(template_seq): aq.append('-'); at.append(template_seq[ti]); ti += 1\n    return ''.join(aq), ''.join(at)\n\n\ndef find_similar_sequences_detailed(query_seq, train_seqs_df, train_coords_dict,\n                                    temporal_cutoff=None, top_n=10):\n    \"\"\"\n    Returns list of tuples: (target_id, train_seq, norm_score, coords, pct_id, aq, at)\n    Sorted by (norm_score, pct_id) descending — original sort order restored.\n\n    Speedup: cached_score() for pre-filter; only run full alignment on survivors.\n    \"\"\"\n    similar_seqs = []\n\n    filtered = (train_seqs_df if temporal_cutoff is None\n                else train_seqs_df[train_seqs_df['temporal_cutoff'] < temporal_cutoff])\n\n    for _, row in filtered.iterrows():\n        target_id, train_seq = row['target_id'], row['sequence']\n        if target_id not in train_coords_dict:\n            continue\n\n        # Fast length pre-filter\n        len_ratio = abs(len(train_seq) - len(query_seq)) / max(len(train_seq), len(query_seq))\n        if len_ratio > 0.4:\n            continue\n\n        # Fast score pre-filter using cache (float, never exhausted)\n        raw_score  = cached_score(query_seq, train_seq)\n        max_score  = 3.0 * min(len(query_seq), len(train_seq))\n        norm_score = raw_score / max(max_score, 1)\n\n        # Skip obviously bad templates early (saves full alignment time)\n        if norm_score < MIN_SIMILARITY * 0.5:\n            continue\n\n        # Full alignment for survivors (fresh object — never cached)\n        alignment  = fresh_alignment(query_seq, train_seq)\n        norm_score = alignment.score / max(max_score, 1)   # recompute from fresh\n\n        identical = total = 0\n        for (qs, qe), (ts, te) in zip(*alignment.aligned):\n            for qp, tp in zip(range(qs, qe), range(ts, te)):\n                total += 1\n                if query_seq[qp] == train_seq[tp]:\n                    identical += 1\n        pct_id = 100.0 * identical / max(len(query_seq), 1)\n\n        aq, at = _build_aligned_strings(query_seq, train_seq, alignment)\n\n        similar_seqs.append((\n            target_id, train_seq, norm_score,\n            train_coords_dict[target_id], pct_id, aq, at\n        ))\n\n    # Original sort order: primary norm_score, secondary pct_id\n    similar_seqs.sort(key=lambda x: (x[2], x[4]), reverse=True)\n    return similar_seqs[:top_n]\n\nprint(\"Sequence search functions ready.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T12:14:40.294552Z","iopub.execute_input":"2026-02-19T12:14:40.29476Z","iopub.status.idle":"2026-02-19T12:14:40.309051Z","shell.execute_reply.started":"2026-02-19T12:14:40.294742Z","shell.execute_reply":"2026-02-19T12:14:40.30815Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ══════════════════════════════════════════════════════════════\n# TEMPLATE ADAPTATION\n# IMPROVEMENT KEPT: Direction-aware gap filling\n#   → estimates local helix direction from flanking atoms\n#   → more physically realistic than fixed [5.95,0,0] vector\n# Uses fresh_alignment() — never cached, never exhausted.\n# ══════════════════════════════════════════════════════════════\n\ndef adapt_template_to_query(query_seq, template_seq, template_coords):\n    \"\"\"Map template C1' coords onto query positions via fresh alignment.\"\"\"\n    # CRITICAL: fresh_alignment() — never a cached object\n    alignment  = fresh_alignment(query_seq, template_seq)\n    new_coords = np.full((len(query_seq), 3), np.nan)\n\n    for (q_start, q_end), (t_start, t_end) in zip(*alignment.aligned):\n        t_chunk = template_coords[t_start:t_end]\n        if len(t_chunk) == (q_end - q_start):\n            new_coords[q_start:q_end] = t_chunk\n\n    # Direction-aware gap filling (improvement over fixed [5.95,0,0])\n    for i in range(len(new_coords)):\n        if not np.isnan(new_coords[i, 0]):\n            continue\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\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:\n            if prev_v > 0 and not np.isnan(new_coords[prev_v-1, 0]):\n                direction = new_coords[prev_v] - new_coords[prev_v-1]\n                d = np.linalg.norm(direction) + 1e-6\n                direction = direction / d * 5.95\n            else:\n                direction = np.array([5.95, 0.0, 0.0])\n            new_coords[i] = new_coords[prev_v] + direction * (i - prev_v)\n        elif next_v >= 0:\n            if next_v < len(new_coords)-1 and not np.isnan(new_coords[next_v+1, 0]):\n                direction = new_coords[next_v+1] - new_coords[next_v]\n                d = np.linalg.norm(direction) + 1e-6\n                direction = direction / d * 5.95\n            else:\n                direction = np.array([5.95, 0.0, 0.0])\n            new_coords[i] = new_coords[next_v] - direction * (next_v - i)\n        else:\n            new_coords[i] = np.array([i * 5.95, 0.0, 0.0])\n\n    return np.nan_to_num(new_coords)\n\n\n# ══════════════════════════════════════════════════════════════\n# CONSTRAINT REFINEMENT\n# IMPROVEMENT KEPT: confidence-adaptive passes\n#   → low confidence (<0.3) gets 1 extra pass for stronger correction\n# ══════════════════════════════════════════════════════════════\n\ndef adaptive_rna_constraints(coordinates, target_id, segments_map,\n                              confidence=1.0, passes=3):\n    coords   = coordinates.copy()\n    segments = segments_map.get(target_id, [(0, len(coords))])\n    strength = 0.85 * (1.0 - min(confidence, 0.95))\n    strength = max(strength, 0.03)\n    effective_passes = passes + (1 if confidence < 0.3 else 0)\n\n    for _ in range(effective_passes):\n        for (s, e) in segments:\n            X = coords[s:e]\n            L = e - s\n            if L < 3:\n                coords[s:e] = X\n                continue\n\n            # Bond-length spring (C1'–C1' ≈ 5.95 Å)\n            d    = X[1:] - X[:-1]\n            dist = np.linalg.norm(d, axis=1) + 1e-6\n            scale = (5.95 - dist) / dist\n            adj   = (d * scale[:, None]) * (0.25 * strength)\n            X[:-1] -= adj;  X[1:] += adj\n\n            # 2-residue span spring (~10.2 Å)\n            if L > 2:\n                d2    = X[2:] - X[:-2]\n                dist2 = np.linalg.norm(d2, axis=1) + 1e-6\n                scale2 = (10.2 - dist2) / dist2\n                adj2   = (d2 * scale2[:, None]) * (0.12 * strength)\n                X[:-2] -= adj2;  X[2:] += adj2\n\n            # Laplacian smoothing (anti-kink)\n            if L > 2:\n                lap = 0.5 * (X[:-2] + X[2:]) - X[1:-1]\n                X[1:-1] += (0.08 * strength) * lap\n\n            # Steric clash repulsion (vectorised)\n            if L >= 25:\n                k   = min(L, 180) if L > 220 else L\n                idx = (np.linspace(0, L-1, k).astype(int) if k < L\n                       else np.arange(L))\n                P    = X[idx]\n                diff = P[:, None, :] - P[None, :, :]\n                distm = np.linalg.norm(diff, axis=2) + 1e-6\n                sep   = np.abs(idx[:, None] - idx[None, :])\n                mask  = (sep > 2) & (distm < 3.3)\n                if np.any(mask):\n                    force = (3.3 - distm) / distm\n                    vec   = (diff * force[:, :, None] * mask[:, :, None]).sum(axis=1)\n                    X[idx] += (0.018 * strength) * vec\n\n            coords[s:e] = X\n\n    return coords\n\nprint(\"Adaptation & refinement functions ready.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T12:14:40.31038Z","iopub.execute_input":"2026-02-19T12:14:40.31068Z","iopub.status.idle":"2026-02-19T12:14:40.327861Z","shell.execute_reply.started":"2026-02-19T12:14:40.310659Z","shell.execute_reply":"2026-02-19T12:14:40.326913Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ══════════════════════════════════════════════════════════════\n# STRUCTURAL DIVERSITY AUGMENTATION\n# FIX #2: Diversity applied PER-TEMPLATE slot, not all-from-one.\n#          Original strategy: 5 different template sources.\n#          Our strategy: same template slots + augmentation within each.\n#\n# apply_prediction_diversity() is called AFTER adapt_template_to_query()\n# for each individual template — so each slot still has a different\n# source structure but also gets an appropriate perturbation for i>0.\n# ══════════════════════════════════════════════════════════════\n\ndef _rotmat(axis, ang):\n    axis = np.asarray(axis, float)\n    axis = axis / (np.linalg.norm(axis) + 1e-12)\n    x, y, z = axis\n    c, s = np.cos(ang), np.sin(ang)\n    C = 1.0 - c\n    return np.array([\n        [c + x*x*C,   x*y*C - z*s, x*z*C + y*s],\n        [y*x*C + z*s, c + y*y*C,   y*z*C - x*s],\n        [z*x*C - y*s, z*y*C + x*s, c + z*z*C]\n    ], dtype=float)\n\ndef apply_hinge(coords, seg, rng, max_angle_deg=20):\n    s, e = seg\n    L = e - s\n    if L < 30: return coords\n    pivot = s + int(rng.integers(10, L - 10))\n    axis  = rng.normal(size=3)\n    ang   = np.deg2rad(float(rng.uniform(-max_angle_deg, max_angle_deg)))\n    R     = _rotmat(axis, ang)\n    X     = coords.copy()\n    p0    = X[pivot].copy()\n    X[pivot+1:e] = (X[pivot+1:e] - p0) @ R.T + p0\n    return X\n\ndef jitter_chains(coords, segments, rng, max_angle_deg=10, max_trans=1.2):\n    X  = coords.copy()\n    gc = X.mean(axis=0, keepdims=True)\n    for (s, e) in segments:\n        axis  = rng.normal(size=3)\n        ang   = np.deg2rad(float(rng.uniform(-max_angle_deg, max_angle_deg)))\n        R     = _rotmat(axis, ang)\n        shift = rng.normal(size=3)\n        shift = shift / (np.linalg.norm(shift) + 1e-12) * float(rng.uniform(0.0, max_trans))\n        c     = X[s:e].mean(axis=0, keepdims=True)\n        X[s:e] = (X[s:e] - c) @ R.T + c + shift\n    X -= X.mean(axis=0, keepdims=True) - gc\n    return X\n\ndef smooth_wiggle(coords, segments, rng, amp=0.6):\n    X = coords.copy()\n    for (s, e) in segments:\n        L = e - s\n        if L < 20: continue\n        ctrl_x   = np.linspace(0, L-1, 6)\n        ctrl_d   = rng.normal(0, amp, size=(6, 3))\n        t        = np.arange(L)\n        disp     = np.vstack([np.interp(t, ctrl_x, ctrl_d[:, k]) for k in range(3)]).T\n        X[s:e]  += disp\n    return X\n\ndef apply_prediction_diversity(adapted, pred_index, target_id, segments, sim, rng):\n    \"\"\"\n    Apply appropriate structural perturbation based on prediction slot index.\n    pred_index=0 → clean template (no perturbation)\n    pred_index=1 → small noise scaled by dissimilarity\n    pred_index=2 → hinge motion on longest domain\n    pred_index=3 → chain rigid-body jitter\n    pred_index=4 → smooth backbone wiggle\n    \"\"\"\n    if pred_index == 0:\n        return adapted.copy()\n    elif pred_index == 1:\n        noise_scale = max(0.01, (0.35 - sim) * 0.08)\n        return adapted + rng.normal(0, noise_scale, adapted.shape)\n    elif pred_index == 2:\n        longest = max(segments, key=lambda se: se[1] - se[0])\n        return apply_hinge(adapted, longest, rng, max_angle_deg=18)\n    elif pred_index == 3:\n        return jitter_chains(adapted, segments, rng, max_angle_deg=8, max_trans=0.8)\n    else:\n        return smooth_wiggle(adapted, segments, rng, amp=0.5)\n\ndef generate_rna_structure(sequence, seed=None):\n    \"\"\"De-novo A-form helix fallback.\"\"\"\n    if seed is not None: np.random.seed(seed)\n    n = len(sequence)\n    coords = np.zeros((n, 3))\n    for i in range(n):\n        angle = i * 0.6\n        coords[i] = [10.0 * np.cos(angle), 10.0 * np.sin(angle), i * 2.5]\n    return coords\n\nprint(\"Diversity augmentation functions ready.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T12:14:40.328665Z","iopub.execute_input":"2026-02-19T12:14:40.328953Z","iopub.status.idle":"2026-02-19T12:14:40.344717Z","shell.execute_reply.started":"2026-02-19T12:14:40.328903Z","shell.execute_reply":"2026-02-19T12:14:40.343891Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from biotite.structure.io.pdbx import CIFFile, get_structure\n\n\ndef extract_c1_atoms(cif_path):\n    cif_file  = CIFFile.read(cif_path)\n    model     = get_structure(cif_file, model=1)\n    chain     = model[model.chain_id == 'A']\n    mask      = chain.atom_name == \"C1'\"\n    c1_atoms  = chain[mask]\n    df        = pd.DataFrame.from_dict(c1_atoms._annot)\n    df['x']   = c1_atoms.coord[:, 0]\n    df['y']   = c1_atoms.coord[:, 1]\n    df['z']   = c1_atoms.coord[:, 2]\n    return df[['res_name', 'res_id', 'x', 'y', 'z']]\n\n\ndef prepare_protenix_json(target_id, sequence, output_path, input_path,\n                           max_length=400):\n    if len(sequence) <= max_length:\n        input_json = [{\n            'sequences': [{\n                'rnaSequence': {\n                    'sequence': sequence,\n                    'count': 1,\n                    'msa': {\n                        'precomputed_msa_dir': f'{input_path}/MSA/{target_id}.MSA.fasta',\n                        'pairing_db': 'rnacentral'\n                    }\n                }\n            }],\n            'name': target_id,\n        }]\n    else:\n        print(f'    Sequence too long ({len(sequence)} > {max_length}), truncating, no MSA')\n        input_json = [{\n            'sequences': [{\n                'rnaSequence': {\n                    'sequence': sequence[:max_length],\n                    'count': 1,\n                }\n            }],\n            'name': target_id,\n        }]\n    json_path = Path(output_path) / 'input_json' / f'{target_id}.json'\n    json_path.parent.mkdir(parents=True, exist_ok=True)\n    with open(json_path, 'w') as f:\n        json.dump(input_json, f, indent=4)\n\n\ndef run_protenix_inference(target_id, sequence, output_path, input_path,\n                            seed=101, n_cycle=12, n_sample=5, n_step=250,\n                            max_length=400):\n    \"\"\"\n    FIX #4 & #6: Single call with n_sample=n_needed, original n_cycle/n_step.\n    No multi-seed loop here — avoids 3x runtime explosion.\n\n    Targeted improvement: short sequences (<80 nt) can afford n_cycle=16\n    since they run fast. Long sequences stay at 12/250.\n    \"\"\"\n    if IS_KAGGLE:\n        checkpoint_path = '/kaggle/input/protenix-finetuned-rna3db-all-1599/1599_ema_0.999.pt'\n    else:\n        checkpoint_path = f'{DATA_PATH}/protenix_chpt/1599_ema_0.999.pt'\n\n    # Targeted quality boost for short sequences only (fast to run)\n    if len(sequence) < 80:\n        n_cycle = 16\n        n_step  = 300\n\n    output_path     = Path(output_path)\n    input_json_path = output_path / 'input_json' / f'{target_id}.json'\n    dump_dir        = output_path / target_id\n    dump_dir.mkdir(parents=True, exist_ok=True)\n    use_msa = 'True' if len(sequence) <= max_length else 'False'\n\n    sys.argv = [\n        'runner/inference.py',\n        f'--seeds={seed}',\n        f'--dump_dir={dump_dir}',\n        f'--input_json_path={input_json_path}',\n        f'--model.N_cycle={n_cycle}',\n        f'--sample_diffusion.N_sample={n_sample}',\n        f'--sample_diffusion.N_step={n_step}',\n        '--augment.use_rnalm True',\n        f'--use_msa {use_msa}',\n        f'--load_checkpoint_path={checkpoint_path}',\n        '',\n    ]\n    from runner.inference import run\n    run()\n\n\ndef get_protenix_predictions(target_id, sequence, output_path,\n                              seed=101, n_sample=5):\n    output_path = Path(output_path)\n    predictions = []\n    for i in range(n_sample):\n        cif_path = (output_path / target_id / target_id /\n                    f'seed_{seed}' / 'predictions' /\n                    f'{target_id}_seed_{seed}_sample_{i}.cif')\n        if cif_path.exists():\n            pred_df = extract_c1_atoms(cif_path)\n            coords  = np.zeros((len(sequence), 3))\n            n_atoms = min(len(pred_df), len(sequence))\n            coords[:n_atoms] = pred_df[['x','y','z']].values[:n_atoms]\n            predictions.append(coords)\n    return predictions\n\n\n@contextlib.contextmanager\ndef protenix_context():\n    original_dir = os.getcwd()\n    os.chdir(PROTENIX_DIR)\n    try:\n        yield\n    finally:\n        os.chdir(original_dir)\n\nprint(\"Protenix interface ready.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T12:14:40.345617Z","iopub.execute_input":"2026-02-19T12:14:40.34598Z","iopub.status.idle":"2026-02-19T12:14:40.584785Z","shell.execute_reply.started":"2026-02-19T12:14:40.345919Z","shell.execute_reply":"2026-02-19T12:14:40.583946Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"template_info_dict       = {}\nprediction_metadata_dict = {}\n\ndef record_template_info(target_id, template_id, similarity, percent_identity):\n    if target_id not in template_info_dict:\n        template_info_dict[target_id] = {\n            'template_ids': [], 'similarities': [], 'percent_identities': []\n        }\n    template_info_dict[target_id]['template_ids'].append(template_id)\n    template_info_dict[target_id]['similarities'].append(similarity)\n    template_info_dict[target_id]['percent_identities'].append(percent_identity)\n\ndef record_prediction_metadata(target_id, pred_num, source,\n                                template_id=None, similarity=None,\n                                percent_identity=None):\n    if target_id not in prediction_metadata_dict:\n        prediction_metadata_dict[target_id] = {}\n    prediction_metadata_dict[target_id][pred_num] = {\n        'source': source,\n        'template_id': template_id if source == 'template' else None,\n        'similarity':  similarity  if source == 'template' else None,\n        'percent_identity': percent_identity if source == 'template' else None,\n    }\n\nprint(\"Metadata tracking ready.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T12:14:40.5857Z","iopub.execute_input":"2026-02-19T12:14:40.586311Z","iopub.status.idle":"2026-02-19T12:14:40.592713Z","shell.execute_reply.started":"2026-02-19T12:14:40.586275Z","shell.execute_reply":"2026-02-19T12:14:40.591937Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ══════════════════════════════════════════════════════════════\n# PREDICT WITH TEMPLATES\n# FIX #2: Restored 1-prediction-per-template strategy.\n#          Each slot i uses a DIFFERENT template as its source.\n#          apply_prediction_diversity() adds slot-appropriate perturbation\n#          ON TOP of the per-template adaptation.\n#\n# FIX #7: Protenix only triggered when templates genuinely insufficient.\n#          No artificial cap at 3 preds to force Protenix on mid-sim targets.\n# ══════════════════════════════════════════════════════════════\n\ndef predict_with_templates(sequence, target_id, train_seqs_df, train_coords_dict,\n                            segments_map, n_predictions=5, temporal_cutoff=None):\n    predictions = []\n    pred_num    = 1\n\n    print('\\n' + '-' * 70)\n    print(f'Target: {target_id} ({len(sequence)} nt)')\n\n    similar_seqs = find_similar_sequences_detailed(\n        sequence, train_seqs_df, train_coords_dict,\n        temporal_cutoff=temporal_cutoff,\n        top_n=n_predictions + 5     # fetch extra so we have fallbacks\n    )\n\n    if similar_seqs:\n        for i, (tmpl_id, tmpl_seq, similarity, tmpl_coords,\n                pct_id, aligned_q, aligned_t) in enumerate(similar_seqs):\n\n            # Stop if below quality threshold\n            if similarity < MIN_SIMILARITY or pct_id < MIN_PERCENT_IDENTITY:\n                if SHOW_ALIGNMENT_DETAILS:\n                    print(f'  Template {i+1}: {tmpl_id} SKIPPED '\n                          f'(sim={similarity:.3f}, id={pct_id:.1f}%) - below threshold')\n                break\n\n            # Reserve 1 slot for Protenix if sequence is short (better quality)\n            if USE_PROTENIX and len(aligned_q) < 100 and i == 4:\n                print(f'  Sequence length {len(aligned_q)} short, leaving 1 space for Protenix')\n                break\n\n            record_template_info(target_id, tmpl_id, similarity, pct_id)\n            record_prediction_metadata(target_id, pred_num, 'template',\n                                       tmpl_id, similarity, pct_id)\n\n            if SHOW_ALIGNMENT_DETAILS:\n                print(f'  Template {i+1}: {tmpl_id} (sim={similarity:.3f}, id={pct_id:.1f}%)')\n\n            # Adapt this template's coordinates to query sequence\n            adapted = adapt_template_to_query(sequence, tmpl_seq, tmpl_coords)\n\n            # Apply slot-appropriate diversity perturbation\n            seed_val  = (abs(hash(target_id)) + i * 10007) % (2**32)\n            rng       = np.random.default_rng(seed_val)\n            segments  = segments_map.get(target_id, [(0, len(sequence))])\n            X         = apply_prediction_diversity(adapted, i, target_id,\n                                                    segments, similarity, rng)\n\n            refined   = adaptive_rna_constraints(X, target_id, segments_map,\n                                                  confidence=similarity, passes=3)\n            predictions.append(refined)\n            pred_num += 1\n\n            if len(predictions) >= n_predictions:\n                break\n\n    n_from_templates = len(predictions)\n    n_needed         = n_predictions - n_from_templates\n\n    if n_needed > 0:\n        print(f'  -> {n_from_templates} from templates, {n_needed} slots for Protenix')\n    else:\n        print(f'  -> All {n_predictions} predictions from templates')\n\n    return predictions, n_needed, pred_num\n\n\ndef generate_predictions_batch(sequences_df, train_seqs_df, train_coords_dict,\n                                dataset_name, use_temporal_cutoff=True,\n                                protenix_output_path=None):\n    start_time = time.time()\n    print(f\"\\n{'='*70}\")\n    print(f'Predicting {len(sequences_df)} {dataset_name} sequences')\n    print(f\"{'='*70}\")\n\n    segments_map, _ = build_segments_map(sequences_df)   # FIX #3: unpack tuple\n\n    print('\\nPHASE 1: Template-based predictions')\n    template_predictions = {}\n    protenix_queue       = {}\n\n    for _, row in sequences_df.iterrows():\n        target_id     = row['target_id']\n        sequence      = row['sequence']\n        temporal_cutoff = (row.get('temporal_cutoff', None)\n                           if use_temporal_cutoff else None)\n\n        preds, n_needed, next_pred = predict_with_templates(\n            sequence, target_id, train_seqs_df, train_coords_dict,\n            segments_map, n_predictions=5, temporal_cutoff=temporal_cutoff\n        )\n        template_predictions[target_id] = preds\n        if n_needed > 0:\n            protenix_queue[target_id] = (n_needed, next_pred, sequence)\n\n    template_time = time.time() - start_time\n    print(f'\\nPhase 1 done: {template_time:.1f}s | '\n          f'{len(protenix_queue)} targets need Protenix')\n\n    protenix_predictions = {}\n\n    if protenix_queue and USE_PROTENIX:\n        print(f'\\nPHASE 2: Protenix for {len(protenix_queue)} targets')\n\n        if protenix_output_path is None:\n            protenix_output_path = Path(OUTPUT_PATH) / f'{dataset_name}_protenix'\n        protenix_output_path = Path(protenix_output_path)\n        protenix_output_path.mkdir(parents=True, exist_ok=True)\n        input_path = Path(DATA_PATH)\n\n        with protenix_context():\n            for i, (target_id, (n_needed, next_pred, sequence)) in                     enumerate(protenix_queue.items()):\n                print(f'\\n  [{i+1}/{len(protenix_queue)}] {target_id} '\n                      f'({len(sequence)} nt, need {n_needed})')\n                try:\n                    prepare_protenix_json(target_id, sequence,\n                                          protenix_output_path, input_path)\n                    t0 = time.time()\n                    # FIX #4: Single call, n_sample=n_needed, seed=101\n                    run_protenix_inference(\n                        target_id, sequence, protenix_output_path, input_path,\n                        seed=101, n_sample=n_needed\n                    )\n                    print(f'    Done in {(time.time()-t0)/60:.1f} min')\n\n                    preds = get_protenix_predictions(\n                        target_id, sequence, protenix_output_path,\n                        seed=101, n_sample=n_needed\n                    )\n                    protenix_predictions[target_id] = preds\n                    print(f'    Got {len(preds)} Protenix predictions')\n\n                except Exception as e:\n                    print(f'    Protenix FAILED: {e}')\n                    print(f'    Will fall back to de novo')\n                    protenix_predictions[target_id] = None\n\n        ptx_time = time.time() - start_time - template_time\n        print(f'\\nPhase 2 done: {ptx_time/60:.1f} min')\n\n    elif protenix_queue and not USE_PROTENIX:\n        print(f'\\nPHASE 2: Protenix disabled, will use de novo for '\n              f'{len(protenix_queue)} targets')\n\n    print('\\nPHASE 3: Combining predictions')\n    all_rows = []\n\n    for _, row in sequences_df.iterrows():\n        target_id   = row['target_id']\n        sequence    = row['sequence']\n        predictions = list(template_predictions[target_id])\n        pred_num    = len(predictions) + 1\n\n        if target_id in protenix_queue:\n            ptx_preds = protenix_predictions.get(target_id)\n            if ptx_preds:\n                for coords in ptx_preds:\n                    record_prediction_metadata(target_id, pred_num, 'protenix')\n                    predictions.append(coords)\n                    pred_num += 1\n                    if len(predictions) >= 5:\n                        break\n\n        n_denovo = 0\n        while len(predictions) < 5:\n            record_prediction_metadata(target_id, pred_num, 'de_novo')\n            seed_val = hash(target_id) % 10000 + len(predictions) * 1000\n            de_novo  = generate_rna_structure(sequence, seed=seed_val)\n            refined  = adaptive_rna_constraints(de_novo, target_id, segments_map,\n                                                 confidence=0.15, passes=3)\n            predictions.append(refined)\n            pred_num += 1\n            n_denovo += 1\n        if n_denovo > 0:\n            print(f'  {target_id}: filled {n_denovo} slots with de novo')\n\n        for j in range(len(sequence)):\n            pred_row = {'ID': f'{target_id}_{j+1}',\n                        'resname': sequence[j], 'resid': j+1}\n            for ii in range(5):\n                pred_row[f'x_{ii+1}'] = predictions[ii][j][0]\n                pred_row[f'y_{ii+1}'] = predictions[ii][j][1]\n                pred_row[f'z_{ii+1}'] = predictions[ii][j][2]\n            all_rows.append(pred_row)\n\n    submission_df = pd.DataFrame(all_rows)\n    col_order = ['ID', 'resname', 'resid']\n    for ii in range(1, 6):\n        for c in ['x', 'y', 'z']:\n            col_order.append(f'{c}_{ii}')\n    submission_df = submission_df[col_order]\n\n    total_time    = time.time() - start_time\n    n_tmpl_only   = sum(1 for tid in sequences_df['target_id']\n                        if tid not in protenix_queue)\n\n    print(f\"\\n{'='*70}\")\n    print(f'{dataset_name.upper()} PREDICTIONS COMPLETE')\n    print(f'  Template-only targets: {n_tmpl_only}')\n    print(f'  Targets with Protenix: {len(protenix_queue)}')\n    print(f'  Total residues:        {len(submission_df)}')\n    print(f'  Runtime:               {total_time:.1f}s ({total_time/60:.1f} min)')\n    print(f\"{'='*70}\\n\")\n    return submission_df\n\nprint(\"Full prediction pipeline ready.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T12:14:40.593794Z","iopub.execute_input":"2026-02-19T12:14:40.594156Z","iopub.status.idle":"2026-02-19T12:14:40.616531Z","shell.execute_reply.started":"2026-02-19T12:14:40.594124Z","shell.execute_reply":"2026-02-19T12:14:40.615749Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"template_info_dict.clear()\nprediction_metadata_dict.clear()\n\ntest_predictions = generate_predictions_batch(\n    test_seqs,\n    combined_seqs,           # train + validation = full template DB\n    combined_coords_dict,\n    dataset_name='test',\n    use_temporal_cutoff=False,\n)\n\ntest_predictions.to_csv('submission.csv', index=False)\nprint('Saved: submission.csv')\ntest_predictions.head(10)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T12:14:40.617532Z","iopub.execute_input":"2026-02-19T12:14:40.617857Z","iopub.status.idle":"2026-02-19T13:09:27.551175Z","shell.execute_reply.started":"2026-02-19T12:14:40.617826Z","shell.execute_reply":"2026-02-19T13:09:27.550339Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Submission format validation ─────────────────────────────\nexpected_cols = (['ID', 'resname', 'resid'] +\n                 [f'{c}_{i}' for i in range(1,6) for c in ['x','y','z']])\nassert list(test_predictions.columns) == expected_cols, 'Column order wrong!'\nassert not test_predictions.isnull().any().any(), 'NaN values found!'\n\n# Check no coordinates collapsed to origin (symptom of cache bug)\nfor i in range(1, 6):\n    coord_cols = [f'x_{i}', f'y_{i}', f'z_{i}']\n    all_zero   = (test_predictions[coord_cols].abs().sum(axis=1) == 0).sum()\n    if all_zero > 0:\n        print(f'WARNING: {all_zero} residues have all-zero coords in prediction {i}')\n    else:\n        print(f'Prediction {i}: coords look healthy (no origin collapse)')\n\n# Source breakdown\nsource_counts = {'template': 0, 'protenix': 0, 'de_novo': 0}\nfor tid, preds in prediction_metadata_dict.items():\n    for pn, meta in preds.items():\n        s = meta.get('source', 'unknown')\n        source_counts[s] = source_counts.get(s, 0) + 1\ntotal_preds = sum(source_counts.values())\nprint(f'\\nPrediction source breakdown:')\nfor src, cnt in source_counts.items():\n    print(f'  {src:<12}: {cnt:>4} ({100*cnt/max(total_preds,1):.1f}%)')\n\nprint(f'\\nSubmission: {len(test_predictions)} rows, {len(test_predictions.columns)} cols')\nprint('All checks PASSED. submission.csv is ready.')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T13:09:27.552219Z","iopub.execute_input":"2026-02-19T13:09:27.55246Z","iopub.status.idle":"2026-02-19T13:09:27.580357Z","shell.execute_reply.started":"2026-02-19T13:09:27.552438Z","shell.execute_reply":"2026-02-19T13:09:27.579464Z"}},"outputs":[],"execution_count":null}]}