{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.12"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":118765,"databundleVersionId":16320058,"isSourceIdPinned":false},{"sourceType":"datasetVersion","sourceId":14874339,"datasetId":9502242,"databundleVersionId":15736806},{"sourceType":"modelInstanceVersion","sourceId":778700,"databundleVersionId":15983433,"modelInstanceId":594204,"modelId":606471,"isSourceIdPinned":false}],"dockerImageVersionId":31259,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":27969.801074,"end_time":"2026-03-26T00:34:28.966371","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2026-03-25T16:48:19.165297","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"7db67a4b","cell_type":"markdown","source":"# Stanford RNA 3D Folding Challenge — Part 2\n\n## Overview\nThis notebook predicts RNA 3D structures (C1′ atom coordinates) from sequences for the Stanford RNA 3D Folding Part 2 Kaggle competition.\n\n**Architecture: RhoFold+-inspired with pre-trained RNA language model**\n\n**Competition Goal:**\n- Predict x, y, z coordinates of the C1′ sugar atom for every residue in each test RNA\n- Submit **5 predictions** per sequence (competition scores the best of 5 by TM-score)\n- Evaluation metric: **TM-score** (0 → random, 1 → perfect; > 0.5 indicates the same fold)\n\n---\n\n## Key Architecture Features\n- **Pre-trained RNA-FM backbone** — `multimolecule/rnafm` (640-dim embeddings, frozen)\n- **Pair representation track** — outer-product mean + triangle attention updates (L × L × d_pair)\n- **Structure module with IPA** — Invariant Point Attention for SE(3)-equivariant coordinate generation\n- **Iterative recycling** — 4 cycles of structure refinement (AlphaFold2-style)\n- **FAPE loss** — Frame Aligned Point Error for rotation/translation-equivariant training\n- **Multi-term loss** — FAPE + dRMSD + soft TM-score + distance distogram\n- **MSA + MI coevolution features** — fed into pair track for secondary structure signal\n- **EMA weight averaging** — Polyak averaging (decay=0.999) for stable inference\n- **Chain-aware PE** — positional encoding resets at chain boundaries\n- **Best-of-5 diversity** — deterministic + MC-dropout + MSA subsampling + retrieval + template\n- **Coordinate noise augmentation** — 0.3 Å Gaussian noise during training\n\n## Dataset Files\n- `train_sequences.csv` + `train_labels.csv` — supervised training\n- `validation_sequences.csv` + `validation_labels.csv` — local TM-score evaluation\n- `test_sequences.csv` — sequences to predict\n- Required output column schema\n- `MSA/` — Multiple Sequence Alignments per target\n- `PDB_RNA/` — Experimental `.cif` files for template features\n- `extra/rna_metadata.csv` — Per-chain PDB quality metrics for training data filtering\n- `extra/parse_fasta_py.py` — Competition-provided FASTA parser\n","metadata":{"papermill":{"duration":0.01022,"end_time":"2026-03-25T16:48:21.916992","exception":false,"start_time":"2026-03-25T16:48:21.906772","status":"completed"},"tags":[]}},{"id":"22156c1f","cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport warnings, re, math, random, copy, os, gc, json, shutil, subprocess\nimport importlib.util\nfrom collections import OrderedDict, Counter\nfrom difflib import SequenceMatcher\nfrom pathlib import Path\nfrom typing import Dict, Tuple, List\nfrom tqdm import tqdm\n\nwarnings.filterwarnings('ignore')\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.checkpoint import checkpoint as ckpt_fn\nfrom torch.utils.data import Dataset, DataLoader\n\ntry:\n    from torch.amp import autocast, GradScaler   # PyTorch >= 2.4\nexcept ImportError:\n    try:\n        from torch.amp import autocast           # PyTorch >= 1.10\n        from torch.cuda.amp import GradScaler    # GradScaler moved later\n    except ImportError:\n        from torch.cuda.amp import autocast, GradScaler\n\ntorch.backends.cudnn.benchmark = True\n\n# RNA-FM uses ESM architecture under the hood\nfrom transformers import EsmModel, EsmConfig\nfrom safetensors.torch import load_file as safetensors_load_file\n\nsns.set_style('whitegrid')\nplt.rcParams['figure.dpi'] = 100\n\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(f\"✓ Libraries loaded  |  Device: {DEVICE}\")\nprint(f\"  PyTorch {torch.__version__}\")","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.execute_input":"2026-03-25T16:48:21.934916Z","iopub.status.busy":"2026-03-25T16:48:21.934203Z","iopub.status.idle":"2026-03-25T16:48:54.833622Z","shell.execute_reply":"2026-03-25T16:48:54.83288Z"},"papermill":{"duration":32.918424,"end_time":"2026-03-25T16:48:54.843323","exception":false,"start_time":"2026-03-25T16:48:21.924899","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"b7dde83c","cell_type":"code","source":"import shutil\nimport subprocess\nfrom pathlib import Path\n\n# ── OFFLINE MODE CONFIGURATION ──\n# By default: tries to install Protenix (if available on Kaggle)\n# To skip Protenix: export OFFLINE_MODE=true\nOFFLINE_MODE = os.environ.get('OFFLINE_MODE', 'false').lower() == 'true'\n\nif OFFLINE_MODE:\n    os.environ['HF_DATASETS_OFFLINE'] = '1'\n    os.environ['TRANSFORMERS_OFFLINE'] = '1'\n    print(\"✓ Offline mode enabled - no internet access required\")\n    print(\"⊘ Skipping optional Protenix installation\")\nelse:\n    # Kaggle-only: copy the best Protenix source tree to writable storage, then install from there.\n    # Dataset: https://www.kaggle.com/datasets/qiweiyin/protenix-v1-adjusted\n    # Contains: Protenix-v1/ and/or Protenix-v1-adjust-v2/ directories\n    protenix_root = Path('/kaggle/input/datasets/qiweiyin/protenix-v1-adjusted')\n    \n    # Try common paths, then fallback to recursive search\n    preferred_setup_paths = [\n        protenix_root / 'Protenix-v1' / 'setup.py',           # Direct subdirectory\n        protenix_root / 'Protenix-v1-adjust-v2' / 'setup.py', # Alternative version\n        protenix_root / 'Protenix-v1-adjust' / 'Protenix-v1' / 'setup.py',\n        protenix_root / 'Protenix-v1-adjust-v2' / 'Protenix-v1-adjust-v2' / 'Protenix-v1' / 'setup.py',\n    ]\n\n    setup_candidates = [path for path in preferred_setup_paths if path.is_file()]\n    if not setup_candidates:\n        # Recursive search if standard paths don't exist\n        setup_candidates = sorted(protenix_root.glob('**/setup.py'))\n\n    print(f'Protenix root: {protenix_root}')\n    print(f'Found {len(setup_candidates)} setup.py files')\n    for candidate in setup_candidates:\n        print(f'  - {candidate}')\n\n    if setup_candidates and protenix_root.exists():\n        setup_path = setup_candidates[0]\n        source_dir = setup_path.parent\n        install_dir = Path('/kaggle/working/protenix_source')\n\n        if install_dir.exists():\n            shutil.rmtree(install_dir)\n        shutil.copytree(source_dir, install_dir)\n\n        print(f'Using best setup: {setup_path}')\n        print(f'Copied source to: {install_dir}')\n\n        try:\n            # Use pip install -e (editable) instead of setup.py install (deprecated)\n            result = subprocess.run(\n                ['pip', 'install', '-e', str(install_dir)],\n                capture_output=True,\n                text=True,\n            )\n            if result.returncode == 0:\n                print(f'Successfully installed protenix from {install_dir}')\n                USE_PROTENIX = True\n            else:\n                print(f'pip install failed with return code {result.returncode}')\n                if result.stderr:\n                    print(f'Error: {result.stderr[:500]}')\n        except Exception as exc:\n            print(f'Install error: {exc}')\n    else:\n        print(f'WARNING: Protenix dataset not found at {protenix_root}')\n        print('This is normal in non-Kaggle environments. Protenix is optional.')","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:48:54.86105Z","iopub.status.busy":"2026-03-25T16:48:54.860234Z","iopub.status.idle":"2026-03-25T16:50:37.002894Z","shell.execute_reply":"2026-03-25T16:50:37.002069Z"},"papermill":{"duration":102.164206,"end_time":"2026-03-25T16:50:37.014988","exception":false,"start_time":"2026-03-25T16:48:54.850782","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"70f740d4","cell_type":"markdown","source":"## 1. Data Loading & Exploration\n\nAvailable files:\n- **train_sequences.csv** + **train_labels.csv** → supervised training (sequence → known 3D coordinates)\n- **validation_sequences.csv** + **validation_labels.csv** → local TM-score evaluation\n- **test_sequences.csv** → sequences to predict\n- Exact required output schema\n- **MSA/** → Multiple Sequence Alignments per target (evolutionary covariation signal)\n- **PDB_RNA/** → Experimental `.cif` files — ground-truth 3D structures used as **templates** for feature extraction\n","metadata":{"papermill":{"duration":0.008728,"end_time":"2026-03-25T16:50:37.032724","exception":false,"start_time":"2026-03-25T16:50:37.023996","status":"completed"},"tags":[]}},{"id":"49705977","cell_type":"code","source":"# Configure paths - Kaggle Competition\nDATASET_NAME = 'stanford-rna-3d-folding-2'\nBASE_PATH    = Path(f'/kaggle/input/{DATASET_NAME}')\nOUTPUT_PATH  = Path('/kaggle/working')\n\n# CSV files\nTRAIN_SEQUENCES      = BASE_PATH / 'train_sequences.csv'\nTRAIN_LABELS         = BASE_PATH / 'train_labels.csv'\nVALIDATION_SEQUENCES = BASE_PATH / 'validation_sequences.csv'\nVALIDATION_LABELS    = BASE_PATH / 'validation_labels.csv'\nTEST_SEQUENCES       = BASE_PATH / 'test_sequences.csv'\n\n# Folders\nMSA_DIR     = BASE_PATH / 'MSA'       # *.MSA.fasta  — evolutionary alignments\nPDB_RNA_DIR = BASE_PATH / 'PDB_RNA'   # *.cif        — experimental 3D structures\nEXTRA_DIR       = BASE_PATH / 'extra'     # quality/structure metadata & helpers\nMETADATA_CSV    = EXTRA_DIR / 'rna_metadata.csv'\nPARSE_FASTA_PY  = EXTRA_DIR / 'parse_fasta_py.py'   # competition-provided FASTA parser\nEXTRA_README    = EXTRA_DIR / 'README.md'            # column descriptions for metadata\n\n# Output\nSUBMISSION_FILE = OUTPUT_PATH / 'submission.csv'\nPROTENIX_CODE_DIR = Path(os.environ.get('PROTENIX_CODE_DIR', '')) if os.environ.get('PROTENIX_CODE_DIR') else None\nPROTENIX_ROOT_DIR = Path(os.environ.get('PROTENIX_ROOT_DIR', str(PROTENIX_CODE_DIR))) if PROTENIX_CODE_DIR else None\nUSE_PROTENIX = bool(PROTENIX_CODE_DIR)\n\nprint(f\"📁 Dataset : {DATASET_NAME}\")\nprint(f\"📂 Input   : {BASE_PATH}\")\nprint(f\"📂 Output  : {OUTPUT_PATH}\")\nprint(f\"📂 MSA     : {MSA_DIR}     (exists={MSA_DIR.exists()})\")\n\nprint(f\"📂 PDB_RNA : {PDB_RNA_DIR} (exists={PDB_RNA_DIR.exists()})\")\nprint(f\"📂 Extra   : {EXTRA_DIR}   (exists={EXTRA_DIR.exists()})\")\nprint(f\"   metadata: {METADATA_CSV.name}  parse_fasta: {PARSE_FASTA_PY.name}  README: {EXTRA_README.name}\")\nprint(f\"\\n✓ Path configuration complete!\")","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:50:37.051477Z","iopub.status.busy":"2026-03-25T16:50:37.051017Z","iopub.status.idle":"2026-03-25T16:50:37.061201Z","shell.execute_reply":"2026-03-25T16:50:37.060499Z"},"papermill":{"duration":0.021105,"end_time":"2026-03-25T16:50:37.062645","exception":false,"start_time":"2026-03-25T16:50:37.04154","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"165878c9","cell_type":"code","source":"# Load all datasets\ntrain_seq    = pd.read_csv(TRAIN_SEQUENCES)\ntrain_labels = pd.read_csv(TRAIN_LABELS)\nval_seq      = pd.read_csv(VALIDATION_SEQUENCES)\nval_labels   = pd.read_csv(VALIDATION_LABELS)\ntest_seq     = pd.read_csv(TEST_SEQUENCES)\n\nprint(\"=\" * 65)\nfor name, df in [(\"train_sequences\",   train_seq),\n                 (\"train_labels\",       train_labels),\n                 (\"val_sequences\",      val_seq),\n                 (\"val_labels\",         val_labels),\n                 (\"test_sequences\",     test_seq)]:\n    print(f\"  {name:<22}  shape={str(df.shape):<16}  cols={df.columns.tolist()}\")\nprint(\"=\" * 65)\n\nprint(\"--- train_sequences (first 3 rows) ---\")\ndisplay(train_seq.head(3))\nprint(\"--- train_labels (first 5 rows) ---\")\ndisplay(train_labels.head(5))\n\n# MSA files\nmsa_files = sorted(MSA_DIR.glob(\"*.fasta\")) if MSA_DIR.exists() else []\nprint(f\" 🧬 MSA: {len(msa_files)} FASTA files\")\nprint(\"  Sample IDs:\", [f.stem.replace('.MSA','') for f in msa_files[:5]], \"...\")\n\n# PDB_RNA .cif files\ncif_files = sorted(PDB_RNA_DIR.glob(\"*.cif\")) if PDB_RNA_DIR.exists() else []\nprint(f\"🗂️  PDB_RNA: {len(cif_files)} CIF files\")\nprint(\"  Sample files:\", [f.name for f in cif_files[:5]], \"...\")","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:50:37.080799Z","iopub.status.busy":"2026-03-25T16:50:37.0805Z","iopub.status.idle":"2026-03-25T16:50:49.477495Z","shell.execute_reply":"2026-03-25T16:50:49.47642Z"},"papermill":{"duration":12.407836,"end_time":"2026-03-25T16:50:49.479163","exception":false,"start_time":"2026-03-25T16:50:37.071327","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"c53dfab7","cell_type":"code","source":"# Infer column names robustly\nSEQ_ID_COL = 'target_id' if 'target_id' in train_seq.columns else train_seq.columns[0]\nSEQ_COL    = 'sequence'  if 'sequence'  in train_seq.columns else train_seq.columns[1]\n\n# Label columns\nLBL_ID_COL = train_labels.columns[0]   # e.g. ID like \"6XRQ_1\"\nprint(f\"Sequence columns : id='{SEQ_ID_COL}'  seq='{SEQ_COL}'\")\nprint(f\"Label columns    : {train_labels.columns.tolist()}\")\n\n# Compute sequence lengths\ntrain_seq['length'] = train_seq[SEQ_COL].str.len()\nval_seq['length']   = val_seq[SEQ_COL].str.len()\ntest_seq['length']  = test_seq[SEQ_COL].str.len()\n\nprint(f\"\\nSequence length statistics:\")\nfor name, df in [(\"train\", train_seq), (\"val\", val_seq), (\"test\", test_seq)]:\n    print(f\"  {name:6s}: n={len(df):4d}  min={df['length'].min():4d}  \"\n          f\"max={df['length'].max():4d}  mean={df['length'].mean():.1f}\")\n\n# Nucleotide composition in training set\nall_nts = ''.join(train_seq[SEQ_COL].values)\nnt_counts = Counter(all_nts)\nprint(f\"\\nTraining nucleotide composition:\")\nfor nt, cnt in sorted(nt_counts.items()):\n    print(f\"  {nt}: {cnt:8d}  ({100*cnt/len(all_nts):.1f}%)\")\n\n# How many labels per sequence?\nprint(f\"\\nTotal training coordinate records : {len(train_labels)}\")\nprint(f\"Expected (sum of seq lengths)    : {train_seq['length'].sum()}\")","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:50:49.499385Z","iopub.status.busy":"2026-03-25T16:50:49.498737Z","iopub.status.idle":"2026-03-25T16:50:50.002504Z","shell.execute_reply":"2026-03-25T16:50:50.001836Z"},"papermill":{"duration":0.515627,"end_time":"2026-03-25T16:50:50.004099","exception":false,"start_time":"2026-03-25T16:50:49.488472","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"35d515b6","cell_type":"code","source":"# Helper to access DataParallel-wrapped models\ndef get_underlying_model(model):\n    \"\"\"Get the underlying model, handling DataParallel wrapping.\"\"\"\n    if isinstance(model, nn.DataParallel):\n        return model.module\n    return model","metadata":{},"outputs":[],"execution_count":null},{"id":"5673b38f","cell_type":"code","source":"# ── Optional Protenix configuration (for advanced users) ──\nUSE_PROTENIX = False      # Use Protenix (requires internet/installation)\nUSE_RNA_MSA = False       # Use RNA MSA features in Protenix\nUSE_TEMPLATE = False      # Use template features\nUSE_MI_FEAT = False       # Use mutual information features\nN_SAMPLE = 1              # Number of diffusion samples for Protenix\nMODEL_NAME = 'rna_esm'    # Protenix model name\n\n# Forward declaration - actual implementation nested inside predict_5_diverse\ndef refine_chain_geometry(coords, lengths, passes=3):\n    \"\"\"Placeholder for refine_chain_geometry - defined inside predict_5_diverse().\"\"\"\n    raise NotImplementedError(\"Call from within predict_5_diverse context\")","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:50:50.024774Z","iopub.status.busy":"2026-03-25T16:50:50.024116Z","iopub.status.idle":"2026-03-25T16:50:51.653148Z","shell.execute_reply":"2026-03-25T16:50:51.652255Z"},"papermill":{"duration":1.641126,"end_time":"2026-03-25T16:50:51.655231","exception":false,"start_time":"2026-03-25T16:50:50.014105","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"31b83be0","cell_type":"markdown","source":"### Metadata-Based Quality Filtering\n\nRemoves noisy training data using `extra/rna_metadata.csv`:\n- Resolution > 9.0 Å removed\n- Fraction observed < 0.7 removed\n- Ribosome subunits excluded (too long, always truncated)\n","metadata":{"papermill":{"duration":0.010125,"end_time":"2026-03-25T16:50:51.676521","exception":false,"start_time":"2026-03-25T16:50:51.666396","status":"completed"},"tags":[]}},{"id":"68fb1f89","cell_type":"code","source":"# Uses extra/rna_metadata.csv to drop noisy training structures.\n# Thresholds are conservative — only removing clearly bad data.\n\nRESOLUTION_MAX   = 9.0    # Å — removes extreme outliers (dataset mean ≈ 3.3)\nFRAC_OBSERVED_MIN = 0.7   # structures with >30% missing residues hurt training\nEXCLUDE_RIBOSOME  = True  # ribosome subunits avg 3,315 nt — always truncated, add noise\n\nn_before = len(train_seq)\n\nif METADATA_CSV.exists():\n    rna_meta = pd.read_csv(METADATA_CSV)\n    print(f\"✓ Loaded rna_metadata.csv: {rna_meta.shape[0]} entries, {rna_meta.shape[1]} columns\")\n\n    # Aggregate per pdb_id (metadata is per-chain, train IDs are per-PDB)\n    rna_meta['_pdb_upper'] = rna_meta['pdb_id'].str.upper()\n    meta_agg = rna_meta.groupby('_pdb_upper').agg(\n        resolution=('resolution', 'first'),\n        fraction_observed=('fraction_observed', 'min'),  # worst chain\n        keyword_ribosome=('keyword_ribosome', 'any'),\n        method=('method', 'first'),\n    ).reset_index()\n\n    train_seq['_pdb_upper'] = train_seq[SEQ_ID_COL].str.upper()\n    train_seq = train_seq.merge(meta_agg, on='_pdb_upper', how='left')\n\n    # Apply filters (keep rows where metadata is missing — be conservative)\n    mask_res  = train_seq['resolution'].isna() | (train_seq['resolution'] <= RESOLUTION_MAX)\n    mask_frac = train_seq['fraction_observed'].isna() | (train_seq['fraction_observed'] >= FRAC_OBSERVED_MIN)\n    mask_ribo = ~train_seq['keyword_ribosome'].fillna(False) if EXCLUDE_RIBOSOME else True\n\n    n_bad_res  = (~mask_res).sum()\n    n_bad_frac = (~mask_frac).sum()\n    n_ribo     = (~mask_ribo).sum() if EXCLUDE_RIBOSOME else 0\n\n    train_seq = train_seq[mask_res & mask_frac & mask_ribo].copy()\n\n    # Clean up temporary columns\n    drop_cols = ['_pdb_upper', 'resolution', 'fraction_observed', 'keyword_ribosome', 'method']\n    train_seq = train_seq.drop(columns=[c for c in drop_cols if c in train_seq.columns])\n\n    n_after = len(train_seq)\n    print(f\"\\n📊 Quality filtering results:\")\n    print(f\"  Resolution > {RESOLUTION_MAX} Å    : {n_bad_res:>5} removed\")\n    print(f\"  Fraction observed < {FRAC_OBSERVED_MIN}  : {n_bad_frac:>5} removed\")\n    print(f\"  Ribosome structures    : {n_ribo:>5} removed\")\n    print(f\"  ──────────────────────────────────\")\n    print(f\"  Kept: {n_after}/{n_before} training sequences ({100*n_after/n_before:.1f}%)\")\n    del rna_meta, meta_agg\n    gc.collect()\nelse:\n    print(\"⚠️  rna_metadata.csv not found — skipping quality filtering\")\n    print(\"  (expected at:\", METADATA_CSV, \")\")\n\n# Pre-filter very long sequences — they always get cropped to MAX_TRAIN_LEN\n# anyway, but build_features + MSA parsing wastes time on full-length data.\nMAX_SEQ_FILTER = 512\nn_pre = len(train_seq)\ntrain_seq = train_seq[train_seq[SEQ_COL].str.len() <= MAX_SEQ_FILTER].reset_index(drop=True)\nprint(f\"  Length filter (≤{MAX_SEQ_FILTER} nt): kept {len(train_seq)}/{n_pre}\")","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:50:51.698054Z","iopub.status.busy":"2026-03-25T16:50:51.697813Z","iopub.status.idle":"2026-03-25T16:50:54.115289Z","shell.execute_reply":"2026-03-25T16:50:54.11413Z"},"papermill":{"duration":2.43042,"end_time":"2026-03-25T16:50:54.117091","exception":false,"start_time":"2026-03-25T16:50:51.686671","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"1a88fdcb","cell_type":"code","source":"def compute_d0(L_ref):\n    \"\"\"Distance scaling factor used in TM-score.\"\"\"\n    if L_ref < 12:  return 0.3\n    if L_ref < 16:  return 0.4\n    if L_ref < 20:  return 0.5\n    if L_ref < 24:  return 0.6\n    if L_ref < 30:  return 0.7\n    return 1.24 * (L_ref - 15) ** (1/3) - 1.8\n\ndef kabsch_align(P, Q):\n    \"\"\"Rotate P to best align with Q (both already mean-centred).\n    Returns rotated P in the centred frame.\n    \"\"\"\n    H = P.T @ Q\n    U, S, Vt = np.linalg.svd(H)\n    d = np.linalg.det(Vt.T @ U.T)\n    D = np.diag([1, 1, d])\n    R = Vt.T @ D @ U.T\n    return P @ R.T\n\ndef tm_score(pred_coords, ref_coords):\n    \"\"\"\n    Compute TM-score between predicted and reference structures.\n    Both inputs: numpy array of shape (L, 3).\n    Residues are matched by index (same numbering).\n    \"\"\"\n    L_ref   = len(ref_coords)\n    L_align = min(len(pred_coords), L_ref)\n    d0      = compute_d0(L_ref)\n\n    # Centre both structures, then Kabsch-rotate P onto Q.\n    # Distances must be computed in the same (centred) frame.\n    P = pred_coords[:L_align].copy()\n    Q = ref_coords[:L_align].copy()\n    P -= P.mean(axis=0)\n    Q -= Q.mean(axis=0)\n    P_aligned = kabsch_align(P, Q)\n\n    dists_sq = np.sum((P_aligned - Q) ** 2, axis=1)\n    score = (1.0 / L_ref) * np.sum(1.0 / (1.0 + dists_sq / d0**2))\n    return float(score)\n\n# Quick sanity check\n_ref = np.random.randn(30, 3) * 5\nprint(f\"TM-score (perfect): {tm_score(_ref, _ref):.4f}  ← should be 1.0\")\nprint(f\"TM-score (random):  {tm_score(np.random.randn(30,3)*5, _ref):.4f}  ← should be low\")\nprint(f\"d0 examples: L=20 → {compute_d0(20):.3f} Å  |  L=100 → {compute_d0(100):.3f} Å\")","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:50:54.141301Z","iopub.status.busy":"2026-03-25T16:50:54.140724Z","iopub.status.idle":"2026-03-25T16:50:54.159384Z","shell.execute_reply":"2026-03-25T16:50:54.158513Z"},"papermill":{"duration":0.031262,"end_time":"2026-03-25T16:50:54.160766","exception":false,"start_time":"2026-03-25T16:50:54.129504","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"7cf4ae68","cell_type":"code","source":"NT_TO_IDX = {'A': 0, 'U': 1, 'G': 2, 'C': 3, 'T': 1}  # T treated as U (DNA→RNA)\n\n# Sequence features\n\ndef one_hot(seq, L=None):\n    \"\"\"Binary nucleotide encoding: each position → [A, U, G, C] indicator vector.\n    Returns (L, 4) float32 array. Unknown nucleotides produce an all-zero row.\n    \"\"\"\n    L   = L or len(seq)\n    mat = np.zeros((L, 4), dtype=np.float32)\n    for i, ch in enumerate(seq[:L]):\n        idx = NT_TO_IDX.get(ch.upper(), -1)\n        if idx >= 0:\n            mat[i, idx] = 1.0\n    return mat\n\ndef positional_encoding(L, d=16):\n    \"\"\"Sinusoidal positional encoding (Vaswani et al. 2017).\n    Returns (L, d) float32 array.\n    \"\"\"\n    pos = np.arange(L)[:, None]\n    i   = np.arange(0, d, 2)[None, :]\n    PE  = np.zeros((L, d), dtype=np.float32)\n    PE[:, 0::2] = np.sin(pos / 10000 ** (i / d))\n    PE[:, 1::2] = np.cos(pos / 10000 ** (i / d))\n    return PE","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:50:54.18266Z","iopub.status.busy":"2026-03-25T16:50:54.182404Z","iopub.status.idle":"2026-03-25T16:50:54.188604Z","shell.execute_reply":"2026-03-25T16:50:54.187837Z"},"papermill":{"duration":0.018622,"end_time":"2026-03-25T16:50:54.189989","exception":false,"start_time":"2026-03-25T16:50:54.171367","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"55129b4a","cell_type":"code","source":"# Chain boundary helpers (multi-chain / stoichiometry-aware)\n\n# Import parse_fasta from the competition-provided extra/parse_fasta_py.py\n_spec = importlib.util.spec_from_file_location(\"parse_fasta_py\", str(PARSE_FASTA_PY))\n_mod  = importlib.util.module_from_spec(_spec)\n_mod.__dict__.update({\"re\": re, \"Dict\": Dict, \"Tuple\": Tuple, \"List\": List})\n_spec.loader.exec_module(_mod)\nparse_fasta = _mod.parse_fasta   # avoid static-analysis 'unresolved import'\nprint(f\"✓ Loaded parse_fasta from {PARSE_FASTA_PY}\")\n\ndef get_chain_lengths(stoichiometry, all_sequences_str, total_len):\n    \"\"\"Return ordered list of (chain_id, length) matching stoichiometry concatenation order.\n\n    Example:\n        stoichiometry = 'A:2;B:1', chain A has 20 nt, chain B has 10 nt\n        → [('A', 20), ('A', 20), ('B', 10)]   (sum = 50 = total_len)\n\n    Falls back to [('?', total_len)] on any parse failure so callers are always safe.\n    \"\"\"\n    try:\n        if not stoichiometry or pd.isna(stoichiometry):\n            raise ValueError(\"missing stoichiometry\")\n        # Use competition-provided parse_fasta to get auth-chain-aware mapping\n        chain_data = (parse_fasta(all_sequences_str)\n                      if all_sequences_str and not pd.isna(all_sequences_str)\n                      else {})\n        # Expand so every auth chain ID (incl. copies) maps to its sequence\n        chain_seqs = {}\n        for _, (seq, all_cids) in chain_data.items():\n            for cid in all_cids:\n                chain_seqs[cid] = seq\n        result = []\n        for part in str(stoichiometry).split(';'):\n            part = part.strip()\n            if ':' in part:\n                chain_id, count = part.split(':')[0].strip(), int(part.split(':')[1])\n            else:\n                chain_id, count = part, 1\n            seq = chain_seqs.get(chain_id, '')\n            if not seq:\n                raise ValueError(f\"chain {chain_id!r} not in all_sequences\")\n            for _ in range(count):\n                result.append((chain_id, len(seq)))\n        if sum(l for _, l in result) != total_len:\n            raise ValueError(\n                f\"length mismatch: chain_sum={sum(l for _, l in result)} != {total_len}\")\n        return result\n    except Exception:\n        return [('?', total_len)]\n\ndef positional_encoding_with_chains(chain_lengths, d=16):\n    \"\"\"Per-chain sinusoidal PE: position counter resets to 0 at each chain boundary.\n\n    For single-chain targets this is identical to positional_encoding(total_L, d).\n    For multi-chain targets chain B's first residue gets position=0, not len(chain_A),\n    so the model cannot confuse 'distance from sequence start' with 'within-chain position'.\n\n    Args:\n        chain_lengths: list of (chain_id, length) from get_chain_lengths()\n        d: PE dimension (must be even)\n    Returns:\n        (total_L, d) float32 array\n    \"\"\"\n    return np.concatenate(\n        [positional_encoding(length, d) for _, length in chain_lengths], axis=0)\n\ndef chain_boundary_feature(chain_lengths):\n    \"\"\"Binary (total_L, 1) feature: 1.0 at the first residue of each chain, else 0.0.\n\n    Gives the model an explicit structural signal about chain boundaries so it\n    can learn that the two 'sides' of a multi-chain complex are separate molecules\n    (not one long single-stranded RNA).\n    \"\"\"\n    total_L = sum(l for _, l in chain_lengths)\n    feat    = np.zeros((total_L, 1), dtype=np.float32)\n    pos = 0\n    for _, length in chain_lengths:\n        feat[pos, 0] = 1.0\n        pos += length\n    return feat","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:50:54.211312Z","iopub.status.busy":"2026-03-25T16:50:54.210965Z","iopub.status.idle":"2026-03-25T16:50:54.227482Z","shell.execute_reply":"2026-03-25T16:50:54.226914Z"},"papermill":{"duration":0.028816,"end_time":"2026-03-25T16:50:54.228923","exception":false,"start_time":"2026-03-25T16:50:54.200107","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"92a6cfc5","cell_type":"code","source":"# MSA features\ndef parse_msa(fasta_path, max_seqs=128):\n    \"\"\"Read a FASTA-formatted MSA file and return a list of sequences.\n\n    Capped at 128 for speed. Sufficient for conservation statistics during training.\n    MI co-evolution uses all available seqs at inference time via explicit max_seqs.\n    \"\"\"\n    seqs, cur = [], ''\n    with open(fasta_path) as f:\n        for line in f:\n            if line.startswith('>'):\n                if cur: seqs.append(cur); cur = ''\n            else:\n                cur += line.strip()\n    if cur: seqs.append(cur)\n    return seqs[:max_seqs]\n\ndef msa_features(fasta_path, seq_len):\n    \"\"\"Per-position nucleotide conservation (4 dims) from MSA.\"\"\"\n    seqs  = parse_msa(fasta_path)\n    if not seqs:\n        return np.zeros((seq_len, 4), dtype=np.float32)\n    freq   = np.zeros((seq_len, 4), dtype=np.float32)\n    counts = np.zeros(seq_len,      dtype=np.float32)\n    for seq in seqs:\n        ungapped = 0\n        for ch in seq:\n            if ch in ('-', '.'):\n                continue\n            if ungapped < seq_len:\n                idx = NT_TO_IDX.get(ch.upper(), -1)\n                if idx >= 0:\n                    freq[ungapped, idx] += 1.0\n                    counts[ungapped]    += 1\n            ungapped += 1\n    np.divide(freq, counts[:, None] + 1e-10, out=freq)\n    return freq\n\ndef msa_mi_features(fasta_path, seq_len, max_L=256):\n    \"\"\"\n    Mutual-information (MI) coevolution features from MSA — 4 dims per position.\n\n    MI(i,j) = sum_ab f(a,b) * log( f(a,b) / (f(a)*f(b)) )\n\n    Source: DeepFoldRNA (Pearce et al. 2022), trRosettaRNA (Wang et al. 2023),\n            reviewed in Li et al. 2020 (Front. Genet. 11:574485).\n    \"\"\"\n    feat = np.zeros((seq_len, 4), dtype=np.float32)\n    if not fasta_path:\n        return feat\n    seqs = parse_msa(fasta_path)\n    if not seqs or len(seqs) < 2:\n        return feat\n\n    L_use = min(seq_len, max_L)\n    N     = len(seqs)\n\n    mat = np.full((N, L_use), 4, dtype=np.int8)\n    for s_idx, seq in enumerate(seqs):\n        ungapped = 0\n        for ch in seq:\n            if ch in ('-', '.'):\n                continue\n            if ungapped >= L_use:\n                break\n            idx = NT_TO_IDX.get(ch.upper(), -1)\n            mat[s_idx, ungapped] = idx if idx >= 0 else 4\n            ungapped += 1\n\n    oh   = (np.arange(5, dtype=np.float32) == mat[:, :, None]).astype(np.float32)\n    fi   = oh.mean(0)\n    ent  = -(fi * np.log(fi + 1e-10)).sum(-1)\n\n    oh_t  = oh.transpose(1, 0, 2)\n    joint = np.tensordot(oh_t, oh_t, ((1,), (1,))) / N\n    fij   = joint.transpose(0, 2, 1, 3)\n\n    fi_i   = fi[:, None, :, None]\n    fi_j   = fi[None, :, None, :]\n    expect = fi_i * fi_j\n\n    log_r     = np.where(fij > 1e-10, np.log((fij + 1e-10) / (expect + 1e-10)), 0.0)\n    mi_matrix = (fij * log_r).sum((-1, -2)).clip(0)\n    np.fill_diagonal(mi_matrix, 0)\n\n    feat[:L_use, 0] = mi_matrix.max(1)\n    feat[:L_use, 1] = mi_matrix.mean(1)\n    feat[:L_use, 2] = mi_matrix.argmax(1) / max(L_use - 1, 1)\n    feat[:L_use, 3] = ent\n    return feat","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:50:54.249871Z","iopub.status.busy":"2026-03-25T16:50:54.249626Z","iopub.status.idle":"2026-03-25T16:50:54.261143Z","shell.execute_reply":"2026-03-25T16:50:54.260329Z"},"papermill":{"duration":0.023431,"end_time":"2026-03-25T16:50:54.26253","exception":false,"start_time":"2026-03-25T16:50:54.239099","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"226aa50f","cell_type":"code","source":"# CIF template features\n\n# Build lookup tables early so template_features() can use cif_lookup\nmsa_lookup = {f.stem.split('.')[0].upper(): f for f in msa_files}\ncif_lookup = {f.stem.upper(): f             for f in cif_files}\nprint(f\"MSA lookup : {len(msa_lookup)} entries\")\nprint(f\"CIF lookup : {len(cif_lookup)} entries\")\n\n_CIF_COORD_CACHE = {}\ndef parse_cif_c1_coords(cif_path):\n    \"\"\"Extract C1' (x,y,z) from mmCIF. Returns {(chain, seq_id) -> np.array}.\"\"\"\n    if cif_path in _CIF_COORD_CACHE:\n        return _CIF_COORD_CACHE[cif_path]\n    coords = {}\n    try:\n        with open(cif_path) as f:\n            for line in f:\n                if not (line.startswith('ATOM') or line.startswith('HETATM')):\n                    continue\n                parts = line.split()\n                if len(parts) < 13 or parts[3].strip('\"') != \"C1'\":\n                    continue\n                chain, seq_id = parts[6], parts[8]\n                try:\n                    coords[(chain, seq_id)] = np.array(\n                        [float(parts[10]), float(parts[11]), float(parts[12])],\n                        dtype=np.float32)\n                except (ValueError, IndexError):\n                    continue\n    except Exception:\n        pass\n    _CIF_COORD_CACHE[cif_path] = coords\n    return coords\n\ndef template_features(pdb_id, seq_len):\n    \"\"\"Mean-centred C1' coords + coverage flag from CIF template. Returns (seq_len, 4).\"\"\"\n    feat     = np.zeros((seq_len, 4), dtype=np.float32)\n    cif_path = cif_lookup.get(str(pdb_id).upper())\n    if cif_path is None:\n        return feat\n    coords_map = parse_cif_c1_coords(cif_path)\n    if not coords_map:\n        return feat\n    xyz_list = []\n    for res_idx in range(1, seq_len + 1):\n        xyz = None\n        for chain in ('A', 'B', 'C', 'D', ' '):\n            xyz = coords_map.get((chain, str(res_idx)))\n            if xyz is not None:\n                break\n        xyz_list.append(xyz)\n    found = [x for x in xyz_list if x is not None]\n    if not found:\n        return feat\n    centroid = np.mean(found, axis=0)\n    for i, xyz in enumerate(xyz_list):\n        if xyz is not None:\n            feat[i, :3] = xyz - centroid\n            feat[i,  3] = 1.0\n    return feat","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:50:54.283908Z","iopub.status.busy":"2026-03-25T16:50:54.283643Z","iopub.status.idle":"2026-03-25T16:50:54.308205Z","shell.execute_reply":"2026-03-25T16:50:54.307446Z"},"papermill":{"duration":0.036791,"end_time":"2026-03-25T16:50:54.309479","exception":false,"start_time":"2026-03-25T16:50:54.272688","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"3346acc5","cell_type":"code","source":"# Robust MSA key resolver + feature sanity checks\n\ndef resolve_msa_key(seq_id):\n    \"\"\"Find MSA file for a seq_id.\n    Tries: (1) exact match, (2) base ID before first '_' (e.g. 'R1107_23' → 'R1107').\n    Returns path or None.\n    \"\"\"\n    key = str(seq_id).upper()\n    if key in msa_lookup:\n        return msa_lookup[key]\n    base = key.split('_')[0]\n    return msa_lookup.get(base, None)\n\n# Quick sanity check\nsample = \"GGCUAGCUAGCUAGCU\"\nL      = len(sample)\noh     = one_hot(sample)\npe     = positional_encoding(L)\n# Chain boundary helpers sanity check (single chain → fallback path)\n_cl    = get_chain_lengths('A:1', None, L)   # triggers fallback → [('?', L)]\nprint(f\"\\nOne-hot  : {oh.shape}\")\nprint(f\"Pos-enc  : {pe.shape}\")\nprint(f\"Chain PE : {positional_encoding_with_chains(_cl).shape}\")\n\nprint(f\"Chain BF : {chain_boundary_feature(_cl).shape}\")\nif cif_files:\n    tf = template_features(cif_files[0].stem.upper(), L)\n    print(f\"CIF tmpl : {tf.shape}  coverage={tf[:, 3].sum():.0f}/{L} residues found\")\nprint(f\"\\n✓ Feature engineering ready\")","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:50:54.330626Z","iopub.status.busy":"2026-03-25T16:50:54.330115Z","iopub.status.idle":"2026-03-25T16:50:54.343722Z","shell.execute_reply":"2026-03-25T16:50:54.343103Z"},"papermill":{"duration":0.025537,"end_time":"2026-03-25T16:50:54.345058","exception":false,"start_time":"2026-03-25T16:50:54.319521","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"825574bf","cell_type":"markdown","source":"## 2. Dataset Preparation\n\n| Feature | Dims | Description |\n|---|---|---|\n| One-hot nucleotide | 4 | A/U/G/C encoding |\n| Chain-aware PE | 16 | Sinusoidal, resets at chain boundaries |\n| Chain boundary | 1 | Binary flag at first residue of each chain |\n| MSA conservation | 4 | Per-position nucleotide frequency |\n| CIF template | 4 | Mean-centred C1' xyz + coverage flag |\n| MI coevolution | 4 | Mutual-information features |\n| **Total** | **33** | `FEATURE_DIM` |\n\nSequences longer than `MAX_TRAIN_LEN=256` are randomly cropped during training. Coordinates are mean-centred.\n","metadata":{"papermill":{"duration":0.009823,"end_time":"2026-03-25T16:50:54.364754","exception":false,"start_time":"2026-03-25T16:50:54.354931","status":"completed"},"tags":[]}},{"id":"aa74dcdc","cell_type":"code","source":"USE_MSA        = True\nUSE_TEMPLATE   = True\nUSE_MI_FEAT    = True\nUSE_CHAIN_BOUNDARY = True   # 1-dim binary feature: 1.0 at first residue of each chain\n\nMAX_TRAIN_LEN = 256   # pair track creates L×L tensors; 256 fits in 15GB T4 with triangle attn\nINFER_WIN_LEN = 384   # window length for long-sequence inference\nINFER_STRIDE  = 192   # 50% overlap blending\n\n# Feature matrix passed to the model: one-hot + chain-aware PE + optional blocks.\n# The learned embedding is fused inside RNATransformer; FEATURE_DIM is the raw input size.\nFEATURE_DIM = (4 + 16\n               + (1 if USE_CHAIN_BOUNDARY else 0)\n               + (4 if USE_MSA           else 0)\n               + (4 if USE_TEMPLATE      else 0)\n               + (4 if USE_MI_FEAT       else 0))\nprint(f\"FEATURE_DIM = {FEATURE_DIM}  \"\n      f\"(4 nt + 16 pos-enc\"\n      f\"{' + 1 chain-boundary' if USE_CHAIN_BOUNDARY else ''}\"\n      f\"{' + 4 MSA' if USE_MSA else ''}\"\n      f\"{' + 4 template' if USE_TEMPLATE else ''}\"\n      f\"{' + 4 MI' if USE_MI_FEAT else ''})\")\nprint(f\"MAX_TRAIN_LEN = {MAX_TRAIN_LEN}  |  INFER_WIN_LEN={INFER_WIN_LEN}, STRIDE={INFER_STRIDE}\")\n\nSEQ_ID_COL  = 'target_id'     if 'target_id'     in train_seq.columns else train_seq.columns[0]\nSEQ_COL     = 'sequence'      if 'sequence'       in train_seq.columns else train_seq.columns[1]\nSTOI_COL    = 'stoichiometry' if 'stoichiometry'  in train_seq.columns else None\nALL_SEQ_COL = 'all_sequences' if 'all_sequences'  in train_seq.columns else None\nprint(f\"Sequence columns: id='{SEQ_ID_COL}', seq='{SEQ_COL}', \"\n      f\"stoichiometry='{STOI_COL}', all_sequences='{ALL_SEQ_COL}'\")\n\n\ndef build_features(seq_id, sequence, stoichiometry=None, all_sequences=None, skip_mi=False):\n    \"\"\"Assemble (L, FEATURE_DIM) feature matrix.\n\n    Uses chain-aware positional encoding (PE resets to 0 at each chain boundary)\n    and an optional 1-dim binary chain-boundary indicator feature so the model can\n    distinguish separate molecular chains in multi-chain targets.\n\n    stoichiometry  — value from the 'stoichiometry'  CSV column (e.g. 'A:2;B:1')\n    all_sequences  — value from the 'all_sequences'  CSV column (FASTA string)\n    \"\"\"\n    L          = len(sequence)\n    chain_lens = get_chain_lengths(stoichiometry, all_sequences, L)\n    parts      = [one_hot(sequence), positional_encoding_with_chains(chain_lens, d=16)]\n    if USE_CHAIN_BOUNDARY:\n        parts.append(chain_boundary_feature(chain_lens))\n    mpath = resolve_msa_key(seq_id)\n    if USE_MSA:\n        parts.append(msa_features(mpath, L) if mpath else np.zeros((L, 4), np.float32))\n    if USE_TEMPLATE:\n        parts.append(template_features(seq_id, L))\n    if USE_MI_FEAT:\n        if skip_mi:\n            parts.append(np.zeros((L, 4), np.float32))\n        else:\n            parts.append(msa_mi_features(mpath, L) if mpath else np.zeros((L, 4), np.float32))\n    return np.concatenate(parts, axis=1).astype(np.float32)\n\n\n# ── Fast label index for O(1) lookup ──────────────────────────────────────\n_LABEL_INDEX_CACHE = {}\n\ndef _get_label_index(labels_df):\n    \"\"\"Build per-prefix row index on first call; cache by DataFrame id.\"\"\"\n    df_key = id(labels_df)\n    if df_key in _LABEL_INDEX_CACHE:\n        return _LABEL_INDEX_CACHE[df_key]\n    if len(labels_df) < 10000:\n        _LABEL_INDEX_CACHE[df_key] = None\n        return None\n    id_col = labels_df.columns[0]\n    lbl_ids = labels_df[id_col].values.astype(str)\n    prefixes = np.array([s.split('_', 1)[0] for s in lbl_ids])\n    idx = {}\n    for pfx in np.unique(prefixes):\n        idx[pfx] = np.where(prefixes == pfx)[0]\n    _LABEL_INDEX_CACHE[df_key] = idx\n    print(f\"  Pre-indexed {len(lbl_ids):,} label rows → {len(idx):,} groups\")\n    return idx\n\n\ndef extract_coords(labels_df, seq_id, stoichiometry=None):\n    \"\"\"Extract ordered C1' (x,y,z) matching the concatenated sequence field.\n\n    CRITICAL — coord ordering must match the sequence field exactly:\n        sequence = chain_A_copy1 + chain_A_copy2 + chain_B_copy1 + ...\n        (concatenated according to stoichiometry, e.g. 'A:2;B:1')\n\n    The labels have per-chain residue numbering (resid restarts at 1 for each\n    chain copy), so sorting by resid ALONE interleaves copies:\n        wrong : copy1_res1, copy2_res1, copy1_res2, copy2_res2, ...\n        correct: copy1_res1..resN, copy2_res1..resN, ...\n\n    Sort priority:\n      1. Chain order from stoichiometry string (when available)\n      2. copy number (1, 2, ...)\n      3. resid within each copy\n    \"\"\"\n    sid_str = str(seq_id)\n    lbl_idx = _get_label_index(labels_df)\n    if lbl_idx is not None and sid_str in lbl_idx:\n        sub = labels_df.iloc[lbl_idx[sid_str]].copy()\n    else:\n        id_col = labels_df.columns[0]\n        mask   = labels_df[id_col].str.startswith(sid_str + '_')\n        sub    = labels_df[mask].copy()\n\n    copy_col  = next((c for c in labels_df.columns if c.lower() == 'copy'),  None)\n    resid_col = next((c for c in labels_df.columns if c.lower() == 'resid'), None)\n    chain_col = next((c for c in labels_df.columns if c.lower() == 'chain'), None)\n\n    sort_cols = []\n    if chain_col and stoichiometry:\n        # Parse \"A:2;B:1\" → ordered chain list [\"A\", \"A\", \"B\"]\n        # Assign rank so chain A < chain B in sort, matching stoichiometry order.\n        chain_order = [part.split(':')[0].strip()\n                       for part in str(stoichiometry).split(';')]\n        chain_rank  = {ch: i for i, ch in enumerate(chain_order)}\n        sub['_chain_rank'] = sub[chain_col].map(chain_rank).fillna(len(chain_order))\n        sort_cols.append('_chain_rank')\n    if copy_col:\n        sort_cols.append(copy_col)\n    if resid_col:\n        sort_cols.append(resid_col)\n\n    if sort_cols:\n        sub = sub.sort_values(by=sort_cols)\n        if '_chain_rank' in sub.columns:\n            sub = sub.drop(columns=['_chain_rank'])\n    else:\n        # Last resort: sort by the third column (original behaviour)\n        sub = sub.sort_values(by=labels_df.columns[2])\n\n    _xyz1      = [c for c in labels_df.columns if c.lower() in ('x_1', 'y_1', 'z_1')]\n    _xyz       = [c for c in labels_df.columns if c.lower() in ('x', 'y', 'z')]\n    coord_cols = _xyz1 or _xyz\n    if len(coord_cols) < 3:\n        coord_cols = labels_df.select_dtypes(include=np.number).columns[-3:].tolist()\n    return sub[coord_cols[:3]].values.astype(np.float32)","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:50:54.38578Z","iopub.status.busy":"2026-03-25T16:50:54.385547Z","iopub.status.idle":"2026-03-25T16:50:54.402263Z","shell.execute_reply":"2026-03-25T16:50:54.401632Z"},"papermill":{"duration":0.029024,"end_time":"2026-03-25T16:50:54.403685","exception":false,"start_time":"2026-03-25T16:50:54.374661","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"094c0397","cell_type":"code","source":"def random_rotation_matrix():\n    \"\"\"Uniform random rotation matrix via QR decomposition of a random normal matrix.\"\"\"\n    M = np.random.randn(3, 3).astype(np.float32)\n    Q, R = np.linalg.qr(M)\n    Q *= np.sign(np.diag(R))   # fix sign so det(Q) = +1\n    if np.linalg.det(Q) < 0:\n        Q[:, 0] *= -1\n    return Q\n\n\nclass RNADataset(Dataset):\n    def __init__(self, seq_df, labels_df, max_len=None, name='dataset', augment=False):\n        self.records = []\n        self.augment = augment\n        id_col    = SEQ_ID_COL\n        total     = len(seq_df)\n        skipped   = 0\n        truncated = 0\n        print(f\"  [{name}] Building {total} sequences ...\")\n        for idx, (_, row) in enumerate(seq_df.iterrows(), 1):\n            sid      = row[id_col]\n            seq      = row[SEQ_COL]\n            stoi     = row[STOI_COL]    if STOI_COL    else None\n            all_seqs = row[ALL_SEQ_COL] if ALL_SEQ_COL else None\n            # Pass stoichiometry so extract_coords can sort by\n            # (chain_stoich_order, copy, resid) — critical for multi-chain targets\n            coords = extract_coords(labels_df, sid, stoichiometry=stoi)\n            if len(coords) != len(seq):\n                skipped += 1\n                continue\n            feat = build_features(sid, seq, stoichiometry=stoi, all_sequences=all_seqs, skip_mi=augment)\n            if max_len and len(seq) > max_len:\n                # Avoid fixed 5' truncation bias: use random crop for train,\n                # centered crop for val/test-like datasets.\n                if augment:\n                    start = np.random.randint(0, len(seq) - max_len + 1)\n                else:\n                    start = max((len(seq) - max_len) // 2, 0)\n                end    = start + max_len\n                feat   = feat[start:end]\n                coords = coords[start:end]\n                seq    = seq[start:end]         # crop sequence to match features\n                truncated += 1\n            coords -= coords.mean(axis=0)\n            self.records.append((feat, coords, sid, seq))\n            if idx % 200 == 0 or idx == total:\n                print(f\"  [{name}] {idx}/{total}  kept={len(self.records)}  \"\n                      f\"skipped={skipped}  truncated={truncated}\")\n        print(f\"  [{name}] Done — {len(self.records)} sequences ready \"\n              f\"(skipped={skipped}, truncated={truncated}, augment={augment})\")\n\n    def __len__(self):\n        return len(self.records)\n\n    def __getitem__(self, i):\n        feat, coords, sid, seq_str = self.records[i]\n        if self.augment:\n            R      = random_rotation_matrix()\n            coords = coords @ R.T\n            # Coordinate noise augmentation (0.3 A std) for generalization\n            coords = coords + np.random.normal(0, 0.3, size=coords.shape).astype(np.float32)\n        return feat, coords, sid, seq_str\n","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:50:54.425403Z","iopub.status.busy":"2026-03-25T16:50:54.425203Z","iopub.status.idle":"2026-03-25T16:50:54.434988Z","shell.execute_reply":"2026-03-25T16:50:54.434387Z"},"papermill":{"duration":0.022181,"end_time":"2026-03-25T16:50:54.436247","exception":false,"start_time":"2026-03-25T16:50:54.414066","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"f55998f7","cell_type":"code","source":"def pad_collate(batch):\n    \"\"\"Pad variable-length sequences → (B, L_max, dim), return boolean mask.\"\"\"\n    feats, coords, ids, seqs = zip(*batch)\n    L_max = max(f.shape[0] for f in feats)\n    B, F, C = len(batch), feats[0].shape[1], 3\n    feat_pad  = torch.zeros(B, L_max, F)\n    coord_pad = torch.zeros(B, L_max, C)\n    mask      = torch.zeros(B, L_max, dtype=torch.bool)\n    for i, (f, c) in enumerate(zip(feats, coords)):\n        L = f.shape[0]\n        feat_pad[i,  :L] = torch.from_numpy(f)\n        coord_pad[i, :L] = torch.from_numpy(c)\n        mask[i,      :L] = True\n    return (feat_pad, coord_pad, mask, list(ids), list(seqs))\n\n\nclass BucketBatchSampler(torch.utils.data.Sampler):\n    \"\"\"Groups sequences of similar length into batches to minimise padding waste.\"\"\"\n    def __init__(self, dataset, batch_size, bucket_size=100, drop_last=False):\n        lengths         = [r[0].shape[0] for r in dataset.records]\n        self.indices    = sorted(range(len(dataset)), key=lambda i: lengths[i])\n        self.batch_size = batch_size\n        self.mega       = batch_size * bucket_size\n        self.drop_last  = drop_last\n\n    def __iter__(self):\n        buckets = [self.indices[i:i + self.mega]\n                   for i in range(0, len(self.indices), self.mega)]\n        all_batches = []\n        for bucket in buckets:\n            shuf = bucket.copy()\n            np.random.shuffle(shuf)\n            for i in range(0, len(shuf), self.batch_size):\n                b = shuf[i:i + self.batch_size]\n                if self.drop_last and len(b) < self.batch_size:\n                    continue\n                all_batches.append(b)\n        np.random.shuffle(all_batches)\n        for batch in all_batches:\n            yield batch\n\n    def __len__(self):\n        n = len(self.indices)\n        return n // self.batch_size if self.drop_last else math.ceil(n / self.batch_size)\n\n\nprint(\"Building datasets...\")\ntrain_dataset = RNADataset(train_seq, train_labels, max_len=MAX_TRAIN_LEN, name='train', augment=True)\nval_dataset   = RNADataset(val_seq,   val_labels,   max_len=MAX_TRAIN_LEN, name='val',   augment=False)\n\nBATCH_SIZE    = 1   # triangle attention is L×L×L memory; batch=1 needed for T4 15GB\ntrain_sampler = BucketBatchSampler(train_dataset, batch_size=BATCH_SIZE, bucket_size=100)\ntrain_loader  = DataLoader(train_dataset, batch_sampler=train_sampler, collate_fn=pad_collate, num_workers=2, pin_memory=True)\nval_loader    = DataLoader(val_dataset,   batch_size=BATCH_SIZE, shuffle=False, collate_fn=pad_collate, num_workers=2, pin_memory=True)\nprint(f\"\\nTrain batches: {len(train_loader)}  |  Val batches: {len(val_loader)}\")\nprint(f\"Effective batch size (with grad accumulation) will be set in training cell\")\nprint(f\"Bucket sampling: ON — sequences grouped by length to reduce padding waste\")\n","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:50:54.457697Z","iopub.status.busy":"2026-03-25T16:50:54.457442Z","iopub.status.idle":"2026-03-25T16:56:04.997917Z","shell.execute_reply":"2026-03-25T16:56:04.997163Z"},"papermill":{"duration":310.553612,"end_time":"2026-03-25T16:56:05.000103","exception":false,"start_time":"2026-03-25T16:50:54.446491","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"8cddaffe","cell_type":"code","source":"print(\"=\" * 60)\nprint(\"DATA PIPELINE DIAGNOSTICS\")\nprint(\"=\" * 60)\n\n# 1. Column names\nprint(f\"\\n[1] train_labels columns : {train_labels.columns.tolist()}\")\nprint(f\"    train_seq    columns : {train_seq.columns.tolist()}\")\nprint(f\"    SEQ_ID_COL={SEQ_ID_COL!r}  SEQ_COL={SEQ_COL!r}  STOI_COL={STOI_COL!r}\")\n\n# 2. A few raw label IDs vs seq IDs\nsample_seq_id = train_seq[SEQ_ID_COL].iloc[0]\nprint(f\"\\n[2] First train seq_id   : {sample_seq_id!r}\")\nmask_ex = train_labels[train_labels.columns[0]].str.startswith(str(sample_seq_id) + '_')\nprint(f\"    Matching label rows  : {mask_ex.sum()}  (expected ≈ seq length)\")\n\n# 2b. Multi-chain / multi-copy audit\ncopy_col_lbl  = next((c for c in train_labels.columns if c.lower() == 'copy'),  None)\nchain_col_lbl = next((c for c in train_labels.columns if c.lower() == 'chain'), None)\nif copy_col_lbl:\n    multi_copy_targets = train_labels[train_labels[copy_col_lbl] > 1][train_labels.columns[0]].str.rsplit('_', n=1).str[0].nunique()\n    print(f\"\\n[2b] copy column found: '{copy_col_lbl}'\")\n    print(f\"    Targets with copy > 1: {multi_copy_targets}  ← these required (copy,resid) sort fix\")\nelse:\n    print(f\"\\n[2b] No 'copy' column found in labels — all targets assumed single-copy\")\nif STOI_COL:\n    multi_chain = (train_seq[STOI_COL].str.contains(';', na=False)).sum()\n    print(f\"    Multi-chain stoichiometry targets: {multi_chain}/{len(train_seq)}\")\n\n# 3. MSA lookup hit rate\nseq_ids  = train_seq[SEQ_ID_COL].astype(str).tolist()\nmsa_hits = sum(1 for sid in seq_ids if resolve_msa_key(sid) is not None)\nprint(f\"\\n[3] MSA lookup hits      : {msa_hits}/{len(seq_ids)} training sequences\")\nif msa_hits == 0:\n    print(\"    ⚠️  NO MSA MATCHES — check msa_lookup key format vs seq IDs\")\n    print(f\"    Sample msa_lookup keys : {list(msa_lookup.keys())[:5]}\")\n    print(f\"    Sample seq IDs         : {seq_ids[:5]}\")\n\n# 4. Dataset load counts\nprint(f\"\\n[4] Train dataset records : {len(train_dataset)}/{len(train_seq)}\")\nprint(f\"    Val   dataset records : {len(val_dataset)}/{len(val_seq)}\")\nif len(train_dataset) < len(train_seq) * 0.5:\n    print(\"    ⚠️  More than 50% of training sequences were SKIPPED — check extract_coords\")\n\n# 5. Spot-check one training sample\nif len(train_dataset) > 0:\n    feat0, coord0, sid0, _ = train_dataset[0]\n    print(f\"\\n[5] Sample[0] id={sid0!r}\")\n    print(f\"    feat  shape={feat0.shape}  (expected L×{FEATURE_DIM})\")\n    print(f\"    coord shape={coord0.shape}  (expected L×3)\")\n    print(f\"    coord mean={coord0.mean(axis=0).round(3)}  (should be ≈ [0,0,0] after centering)\")\n    print(f\"    coord std ={coord0.std(axis=0).round(3)}   (typical RNA: 10–80 Å)\")\n    # Sanity: consecutive C1' distances should be 5-7 Å in correct order\n    diffs = np.linalg.norm(np.diff(coord0, axis=0), axis=1)\n    pct_ok = (diffs < 15).mean() * 100\n    print(f\"    consecutive C1' dist: mean={diffs.mean():.1f} Å  ({pct_ok:.0f}% < 15 Å) \"\n          f\"← should be ~5-7 Å if coord order is correct\")\n\nprint(\"\\n\" + \"=\" * 60)\n","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:56:05.04038Z","iopub.status.busy":"2026-03-25T16:56:05.040072Z","iopub.status.idle":"2026-03-25T16:56:07.529208Z","shell.execute_reply":"2026-03-25T16:56:07.528249Z"},"papermill":{"duration":2.510798,"end_time":"2026-03-25T16:56:07.530839","exception":false,"start_time":"2026-03-25T16:56:05.020041","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"bc448609","cell_type":"code","source":"# Build a lightweight library from labeled train+validation structures.\n# For each test sequence we can retrieve similar sequences and reuse their geometry\n# as an extra candidate. Since Kaggle scores best-of-5, this can improve robustness.\n\ndef _coords_are_valid(coords, max_abs=1e5):\n    c = np.asarray(coords, dtype=np.float32)\n    return np.isfinite(c).all() and (np.abs(c) <= max_abs).all()\n\n\ndef _resample_coords_linear(coords, target_len):\n    \"\"\"Resample (L,3) coordinates to target_len with linear interpolation.\"\"\"\n    c = np.asarray(coords, dtype=np.float32)\n    if not _coords_are_valid(c):\n        return None\n    if len(c) == target_len:\n        out = c.copy()\n    elif len(c) < 2 or target_len < 2:\n        out = np.tile(c[:1], (target_len, 1)).astype(np.float32)\n    else:\n        x_old = np.linspace(0.0, 1.0, len(c), dtype=np.float32)\n        x_new = np.linspace(0.0, 1.0, target_len, dtype=np.float32)\n        out = np.stack([np.interp(x_new, x_old, c[:, d]) for d in range(3)], axis=1).astype(np.float32)\n    if not np.isfinite(out).all():\n        return None\n    out -= out.mean(axis=0, keepdims=True)\n    if not np.isfinite(out).all() or (np.abs(out) > 1e5).any():\n        return None\n    return out\n\n\ndef _seq_similarity(a, b):\n    \"\"\"Fast normalized sequence similarity in [0,1].\"\"\"\n    if not a or not b:\n        return 0.0\n    return float(SequenceMatcher(None, a, b).ratio())\n\n\nprint(\"Building retrieval bank from dataset records (no re-extraction) ...\")\nRETRIEVAL_DB = []\nfor ds, tag in [(train_dataset, 'train'), (val_dataset, 'val')]:\n    kept = 0\n    skipped_bad = 0\n    for feat, coords, sid, seq_str in ds.records:\n        if len(coords) < 8:\n            continue\n        c = coords.astype(np.float32).copy()\n        if not _coords_are_valid(c):\n            skipped_bad += 1\n            continue\n        c -= c.mean(axis=0, keepdims=True)\n        if not np.isfinite(c).all() or (np.abs(c) > 1e5).any():\n            skipped_bad += 1\n            continue\n        RETRIEVAL_DB.append({'id': sid, 'seq': seq_str, 'coords': c})\n        kept += 1\n    print(f\"  {tag:5s}: {kept} structures added (bad skipped={skipped_bad})\")\n\nprint(f\"✓ Retrieval bank size: {len(RETRIEVAL_DB)}\")\n","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:56:07.555431Z","iopub.status.busy":"2026-03-25T16:56:07.555065Z","iopub.status.idle":"2026-03-25T16:56:07.670842Z","shell.execute_reply":"2026-03-25T16:56:07.669915Z"},"papermill":{"duration":0.129463,"end_time":"2026-03-25T16:56:07.672417","exception":false,"start_time":"2026-03-25T16:56:07.542954","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"b9b1e745","cell_type":"code","source":"def _kmer_similarity(a, b, k=3):\n    \"\"\"Fast k-mer Jaccard similarity — O(L) instead of O(L^2) SequenceMatcher.\"\"\"\n    if len(a) < k or len(b) < k:\n        return _seq_similarity(a, b)\n    sa = set(a[i:i+k] for i in range(len(a) - k + 1))\n    sb = set(b[i:i+k] for i in range(len(b) - k + 1))\n    inter = len(sa & sb)\n    union = len(sa | sb)\n    return inter / max(union, 1)\n\n\ndef retrieval_candidate(seq, top_k=3, min_score=0.35):\n    \"\"\"Return a retrieved coordinate candidate for `seq`, or None if weak match.\n    Uses k-mer similarity for speed, then refines top hits with SequenceMatcher.\n    Lowered threshold (0.35) and increased top_k (3) for better best-of-5 coverage.\"\"\"\n    if not RETRIEVAL_DB:\n        return None\n\n    scored = []\n    for item in RETRIEVAL_DB:\n        s = _kmer_similarity(seq, item['seq'])\n        if s >= min_score * 0.8:\n            scored.append((s, item))\n    if not scored:\n        return None\n\n    scored.sort(key=lambda x: x[0], reverse=True)\n    refined = []\n    for _, item in scored[:min(10, len(scored))]:\n        s = _seq_similarity(seq, item['seq'])\n        if s >= min_score:\n            refined.append((s, item))\n    if not refined:\n        return None\n\n    refined.sort(key=lambda x: x[0], reverse=True)\n    top = refined[:top_k]\n\n    preds = []\n    wts = []\n    for s, item in top:\n        pred = _resample_coords_linear(item['coords'], len(seq))\n        if pred is None:\n            continue\n        preds.append(pred)\n        wts.append(max(s, 1e-4))\n\n    if not preds:\n        return None\n\n    w = np.array(wts, dtype=np.float32)\n    w /= w.sum()\n    out = np.zeros((len(seq), 3), dtype=np.float32)\n    for wi, pi in zip(w, preds):\n        out += wi * pi\n    out -= out.mean(axis=0, keepdims=True)\n    if not np.isfinite(out).all() or (np.abs(out) > 1e5).any():\n        return None\n    return out\n\n\ndef template_candidate(seq_id, seq_len, min_coverage_frac=0.20):\n    \"\"\"Use CIF template xyz directly when enough residues are covered.\"\"\"\n    tf = template_features(seq_id, seq_len)\n    cov = tf[:, 3] > 0.5\n    if cov.sum() < max(6, int(min_coverage_frac * seq_len)):\n        return None\n    xyz = tf[:, :3].astype(np.float32)\n    cov_mask = tf[:, 3] > 0.5\n    if cov_mask.any():\n        xyz -= xyz[cov_mask].mean(axis=0, keepdims=True)\n    return xyz\n","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:56:07.696977Z","iopub.status.busy":"2026-03-25T16:56:07.696693Z","iopub.status.idle":"2026-03-25T16:56:07.707382Z","shell.execute_reply":"2026-03-25T16:56:07.706608Z"},"papermill":{"duration":0.024408,"end_time":"2026-03-25T16:56:07.708882","exception":false,"start_time":"2026-03-25T16:56:07.684474","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"c3b66a82","cell_type":"markdown","source":"## 3. Model Architecture — Pre-trained RNA-FM + Pair Track + IPA Structure Module\n\n**Single-representation track (per-residue):**\n- Pre-trained RNA-FM (640-dim, frozen) → project to `d_model=256`\n- Concatenate MSA/MI/template/chain features (33-dim) → fuse via linear projection\n- 4 Transformer encoder layers with RPB for sequence-level context\n\n**Pair representation track (per-residue-pair):**\n- Initialized from outer-product mean of single representations\n- 2 rounds of triangle attention (starting node + ending node) + triangle multiplication\n- Captures inter-residue distance/orientation relationships\n\n**Structure module:**\n- 4-layer IPA (Invariant Point Attention) with 3D query/key point attention\n- Per-layer coordinate update: delta xyz from single representation → C1′ coordinates\n- 4 recycling iterations: predicted coords fed back to refine both tracks\n\n**Loss function:**\n| Term | Weight | Description |\n|---|---|---|\n| FAPE (local) | 30% | Frame-aligned point error within 10 Å clamp |\n| dRMSD | 25% | Distance-based RMSD (12 Å cutoff) |\n| Soft TM-score | 25% | Differentiable TM-score (warmup after epoch 10) |\n| Distogram CE | 20% | Cross-entropy on binned inter-residue distances |","metadata":{"papermill":{"duration":0.010817,"end_time":"2026-03-25T16:56:07.730803","exception":false,"start_time":"2026-03-25T16:56:07.719986","status":"completed"},"tags":[]}},{"id":"06b9a348","cell_type":"code","source":"# Pre-trained RNA Language Model — feature extractor (frozen)\n# Source: multimolecule/rnafm on HuggingFace, uploaded as Kaggle Model input\n# RNA-FM: 12-layer BERT (ESM architecture), 640-dim, 99.5M params\n\nRNA_FM_DIM = 640  # hidden size\nRNA_FM_PATH = \"/kaggle/input/models/emanafi/multimolecule-rnafm/transformers/transformers/1\"\n\nclass SimpleRNATokenizer:\n    \"\"\"RNA tokenizer — reads vocab.txt from the model directory.\"\"\"\n    def __init__(self, model_path):\n        vocab_path = Path(model_path) / 'vocab.txt'\n        self.vocab = {}\n        with open(vocab_path) as f:\n            for idx, line in enumerate(f):\n                self.vocab[line.strip()] = idx\n\n        tok_cfg_path = Path(model_path) / 'tokenizer_config.json'\n        if tok_cfg_path.exists():\n            with open(tok_cfg_path) as f:\n                tok_cfg = json.load(f)\n            self.replace_T = tok_cfg.get('replace_T_with_U', True)\n        else:\n            self.replace_T = True\n\n        self.cls_id = self.vocab.get('<cls>', 1)\n        self.eos_id = self.vocab.get('<eos>', 2)\n        self.unk_id = self.vocab.get('<unk>', 3)\n        self.pad_id = self.vocab.get('<pad>', 0)\n        print(f\"  Tokenizer: {len(self.vocab)} tokens, pad={self.pad_id}, cls={self.cls_id}, eos={self.eos_id}\")\n\n    def __call__(self, sequences, max_length=1026):\n        all_ids = []\n        for seq in sequences:\n            seq = seq.upper().replace(' ', '')\n            if self.replace_T:\n                seq = seq.replace('T', 'U')\n            ids = [self.cls_id]\n            for ch in seq[:max_length - 2]:\n                ids.append(self.vocab.get(ch, self.unk_id))\n            ids.append(self.eos_id)\n            all_ids.append(ids)\n        max_len = max(len(x) for x in all_ids)\n        padded = [x + [self.pad_id] * (max_len - len(x)) for x in all_ids]\n        masks = [[1] * len(x) + [0] * (max_len - len(x)) for x in all_ids]\n        return {'input_ids': torch.tensor(padded, dtype=torch.long),\n                'attention_mask': torch.tensor(masks, dtype=torch.long)}\n\n\ndef _map_rnafm_key(k):\n    \"\"\"Map multimolecule RNA-FM weight key → EsmModel key.\n    Handles: model. prefix, layer_norm→LayerNorm casing, embedding/encoder LN positions.\"\"\"\n    # Skip pretraining heads (lm_head, ss_head)\n    if k.startswith(('lm_head.', 'ss_head.')):\n        return None\n    # Strip 'model.' prefix\n    if k.startswith('model.'):\n        k = k[len('model.'):]\n    # Post-encoder layer norm → ESM's emb_layer_norm_after\n    if k.startswith('encoder.layer_norm.'):\n        return k.replace('encoder.layer_norm.', 'encoder.emb_layer_norm_after.')\n    # Pre-encoder embedding layer norm — no ESM equivalent, skip\n    if k.startswith('embeddings.layer_norm.'):\n        return None\n    # Per-layer attention layer_norm → LayerNorm\n    k = k.replace('.attention.layer_norm.', '.attention.LayerNorm.')\n    # Per-layer output layer_norm → LayerNorm\n    k = re.sub(r'\\.layer\\.(\\d+)\\.layer_norm\\.', r'.layer.\\1.LayerNorm.', k)\n    return k\n\n\ndef _load_rnafm_as_esm(model_path):\n    \"\"\"Load RNA-FM weights into an EsmModel. Reads config.json for all hyperparameters.\"\"\"\n    with open(Path(model_path) / 'config.json') as f:\n        raw_cfg = json.load(f)\n    esm_cfg = EsmConfig(\n        vocab_size=raw_cfg.get('vocab_size', 26),\n        hidden_size=raw_cfg.get('hidden_size', 640),\n        num_hidden_layers=raw_cfg.get('num_hidden_layers', 12),\n        num_attention_heads=raw_cfg.get('num_attention_heads', 20),\n        intermediate_size=raw_cfg.get('intermediate_size', 5120),\n        max_position_embeddings=raw_cfg.get('max_position_embeddings', 1026),\n        hidden_dropout_prob=raw_cfg.get('hidden_dropout', 0.1),\n        attention_probs_dropout_prob=raw_cfg.get('attention_dropout', 0.1),\n        pad_token_id=raw_cfg.get('pad_token_id', 0),\n        layer_norm_eps=raw_cfg.get('layer_norm_eps', 1e-12),\n    )\n    esm_model = EsmModel(esm_cfg)\n\n    st_path = Path(model_path) / 'model.safetensors'\n    bin_path = Path(model_path) / 'pytorch_model.bin'\n    if st_path.exists():\n        state = safetensors_load_file(str(st_path))\n    elif bin_path.exists():\n        state = torch.load(str(bin_path), map_location='cpu', weights_only=True)\n    else:\n        raise FileNotFoundError(f'No model weights found in {model_path}')\n\n    # Map multimolecule key names to EsmModel key names\n    mapped = {}\n    for k, v in state.items():\n        new_k = _map_rnafm_key(k)\n        if new_k is not None:\n            mapped[new_k] = v\n\n    missing, unexpected = esm_model.load_state_dict(mapped, strict=False)\n    # Only warn about keys that aren't contact_head or pooler (both unused for embeddings)\n    real_missing = [k for k in missing if 'contact_head' not in k and 'pooler' not in k]\n    if real_missing:\n        print(f'  Warning: {len(real_missing)} missing keys (first 5): {real_missing[:5]}')\n    if unexpected:\n        print(f'  Warning: {len(unexpected)} unexpected keys (first 5): {unexpected[:5]}')\n    return esm_model\n\n\nclass RNAFMEmbedder(nn.Module):\n    \"\"\"Frozen pre-trained RNA-FM for per-residue embeddings.\n    Includes a CPU cache for repeated sequences (avoids 12-layer forward).\"\"\"\n    def __init__(self, model_path=RNA_FM_PATH, max_len=1022):\n        super().__init__()\n        self.tokenizer = SimpleRNATokenizer(model_path)\n        self.model = _load_rnafm_as_esm(model_path)\n        self.max_len = max_len\n        self._cache = OrderedDict()\n        self._cache_maxsize = 512\n        for p in self.model.parameters():\n            p.requires_grad = False\n        self.model.eval()\n        n_params = sum(p.numel() for p in self.model.parameters())\n        print(f\"✓ RNA-FM loaded from {model_path}: {n_params/1e6:.1f}M parameters (frozen)\")\n\n    @torch.no_grad()\n    def forward(self, sequences: list) -> torch.Tensor:\n        # Check cache — for batch_size=1, this hits every training step after epoch 1\n        if len(sequences) == 1 and sequences[0] in self._cache:\n            return self._cache[sequences[0]].to(next(self.model.parameters()).device)\n\n        seq_lens = [len(s[:self.max_len]) for s in sequences]\n        enc = self.tokenizer(sequences, max_length=self.max_len + 2)\n        enc = {k: v.to(next(self.model.parameters()).device) for k, v in enc.items()}\n        out = self.model(**enc)\n        emb = out.last_hidden_state\n\n        B = emb.shape[0]\n        max_seq_len = max(seq_lens)\n        device = emb.device\n        result = torch.zeros(B, max_seq_len, emb.shape[-1], device=device, dtype=emb.dtype)\n        for i, L in enumerate(seq_lens):\n            result[i, :L] = emb[i, 1:1+L]\n\n        # Cache on CPU for batch_size=1 (LRU-bounded)\n        if B == 1:\n            key = sequences[0]\n            if key in self._cache:\n                self._cache.move_to_end(key)\n            self._cache[key] = result.cpu()\n            if len(self._cache) > self._cache_maxsize:\n                self._cache.popitem(last=False)\n        return result\n\n# Load RNA-FM\ntry:\n    rna_fm = RNAFMEmbedder().to(DEVICE)\nexcept Exception as e:\n    if OFFLINE_MODE:\n        print(f\"⊘ Offline mode: Could not load RNA-FM model locally\")\n        print(f\"  Error: {e}\")\n        print(f\"  RNA-FM model must be pre-downloaded to: {RNA_FM_PATH}\")\n        rna_fm = None\n    else:\n        raise\n\n# Quick test\n_test_emb = rna_fm([\"AUGCAUGC\"])\nprint(f\"  Test embedding: input='AUGCAUGC' -> shape={_test_emb.shape}\")\nprint(f\"  Embedding dim: {_test_emb.shape[-1]}\")","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:56:07.754311Z","iopub.status.busy":"2026-03-25T16:56:07.754026Z","iopub.status.idle":"2026-03-25T16:56:13.676745Z","shell.execute_reply":"2026-03-25T16:56:13.675853Z"},"papermill":{"duration":5.936566,"end_time":"2026-03-25T16:56:13.678251","exception":false,"start_time":"2026-03-25T16:56:07.741685","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"44a3997b","cell_type":"code","source":"# Pair Representation Track — outer product mean + triangle attention\n\nD_MODEL = 256\nD_PAIR  = 64\nN_HEAD  = 8\nN_SINGLE_LAYERS = 4\nN_PAIR_LAYERS   = 2\nN_IPA_LAYERS    = 4\nDIM_FF  = 1024\nDROPOUT = 0.1\nCOORD_CLAMP = 200.0\n\n\nclass OuterProductMean(nn.Module):\n    \"\"\"Outer product mean: single → pair representation (AlphaFold2 Evoformer).\"\"\"\n    def __init__(self, d_single, d_pair, d_hidden=32):\n        super().__init__()\n        self.norm = nn.LayerNorm(d_single)\n        self.proj_a = nn.Linear(d_single, d_hidden)\n        self.proj_b = nn.Linear(d_single, d_hidden)\n        self.out = nn.Linear(d_hidden * d_hidden, d_pair)\n\n    def forward(self, s, mask=None):\n        \"\"\"s: (B, L, d_single) → (B, L, L, d_pair)\"\"\"\n        s = self.norm(s)\n        a = self.proj_a(s)  # (B, L, h)\n        b = self.proj_b(s)  # (B, L, h)\n        B, L, h = a.shape\n        # Outer product: a_i ⊗ b_j → (B, L, L, h*h)\n        outer = torch.einsum('bih,bjk->bijhk', a, b).reshape(B, L, L, h * h)\n        if mask is not None:\n            pair_mask = mask.unsqueeze(-1) & mask.unsqueeze(-2)  # (B, L, L)\n            outer = outer * pair_mask.unsqueeze(-1).float()\n        return self.out(outer)\n\n\nclass TriangleAttention(nn.Module):\n    \"\"\"Triangle self-attention around starting/ending node.\"\"\"\n    def __init__(self, d_pair, n_head=4, starting=True):\n        super().__init__()\n        self.starting = starting\n        self.norm = nn.LayerNorm(d_pair)\n        self.mha = nn.MultiheadAttention(d_pair, n_head, dropout=0.0, batch_first=True)\n        self.gate = nn.Sequential(nn.Linear(d_pair, d_pair), nn.Sigmoid())\n\n    def forward(self, z):\n        \"\"\"z: (B, L, L, d_pair)\"\"\"\n        B, L, _, D = z.shape\n        z_in = self.norm(z)\n        g = self.gate(z_in)\n        if self.starting:\n            z_flat = z_in.reshape(B * L, L, D)\n            z_att, _ = self.mha(z_flat, z_flat, z_flat)\n            z_att = z_att.reshape(B, L, L, D)\n        else:\n            z_t = z_in.permute(0, 2, 1, 3).reshape(B * L, L, D)\n            z_att, _ = self.mha(z_t, z_t, z_t)\n            z_att = z_att.reshape(B, L, L, D).permute(0, 2, 1, 3)\n        return z + g * z_att\n\n\nclass TriangleMultiplication(nn.Module):\n    \"\"\"Triangle multiplicative update (outgoing / incoming).\"\"\"\n    def __init__(self, d_pair, d_hidden=64, outgoing=True):\n        super().__init__()\n        self.outgoing = outgoing\n        self.norm = nn.LayerNorm(d_pair)\n        self.proj_a = nn.Linear(d_pair, d_hidden)\n        self.proj_b = nn.Linear(d_pair, d_hidden)\n        self.gate_a = nn.Sequential(nn.Linear(d_pair, d_hidden), nn.Sigmoid())\n        self.gate_b = nn.Sequential(nn.Linear(d_pair, d_hidden), nn.Sigmoid())\n        self.out_norm = nn.LayerNorm(d_hidden)\n        self.out_proj = nn.Linear(d_hidden, d_pair)\n        self.out_gate = nn.Sequential(nn.Linear(d_pair, d_pair), nn.Sigmoid())\n\n    def forward(self, z):\n        z_in = self.norm(z)\n        a = self.proj_a(z_in) * self.gate_a(z_in)\n        b = self.proj_b(z_in) * self.gate_b(z_in)\n        if self.outgoing:\n            update = torch.einsum('bikh,bjkh->bijh', a, b)\n        else:\n            update = torch.einsum('bkih,bkjh->bijh', a, b)\n        update = self.out_proj(self.out_norm(update))\n        return z + self.out_gate(z_in) * update\n\n\nclass PairTrack(nn.Module):\n    \"\"\"Full pair representation track: OPM init + triangle updates.\"\"\"\n    def __init__(self, d_single=D_MODEL, d_pair=D_PAIR, n_layers=N_PAIR_LAYERS):\n        super().__init__()\n        self.opm = OuterProductMean(d_single, d_pair)\n        self.layers = nn.ModuleList()\n        for _ in range(n_layers):\n            self.layers.append(nn.ModuleDict({\n                'tri_att_start': TriangleAttention(d_pair, n_head=4, starting=True),\n                'tri_att_end':   TriangleAttention(d_pair, n_head=4, starting=False),\n                'tri_mul_out':   TriangleMultiplication(d_pair, outgoing=True),\n                'tri_mul_in':    TriangleMultiplication(d_pair, outgoing=False),\n                'pair_ff':       nn.Sequential(\n                    nn.LayerNorm(d_pair), nn.Linear(d_pair, d_pair * 4),\n                    nn.GELU(), nn.Linear(d_pair * 4, d_pair)),\n                'dropout':       nn.Dropout(DROPOUT),\n            }))\n\n    def _pair_layer_step(self, layer, z):\n        \"\"\"One pair layer — wrapped for gradient checkpointing.\"\"\"\n        z = layer['tri_att_start'](z)\n        z = layer['dropout'](z)\n        z = layer['tri_att_end'](z)\n        z = layer['dropout'](z)\n        z = layer['tri_mul_out'](z)\n        z = layer['dropout'](z)\n        z = layer['tri_mul_in'](z)\n        z = layer['dropout'](z)\n        z = z + layer['pair_ff'](z)\n        return z\n\n    def forward(self, single, pair_init=None, mask=None):\n        z = self.opm(single, mask)\n        if pair_init is not None:\n            z = z + pair_init\n        for layer in self.layers:\n            if self.training and z.requires_grad:\n                z = ckpt_fn(self._pair_layer_step, layer, z, use_reentrant=False)\n            else:\n                z = self._pair_layer_step(layer, z)\n        return z\n\nprint(f\"✓ Pair track: d_pair={D_PAIR}, {N_PAIR_LAYERS} triangle rounds\")","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:56:13.704874Z","iopub.status.busy":"2026-03-25T16:56:13.70438Z","iopub.status.idle":"2026-03-25T16:56:13.723521Z","shell.execute_reply":"2026-03-25T16:56:13.722769Z"},"papermill":{"duration":0.034147,"end_time":"2026-03-25T16:56:13.724949","exception":false,"start_time":"2026-03-25T16:56:13.690802","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"577caf1f","cell_type":"code","source":"# Invariant Point Attention (IPA) — Structure Module\n# AlphaFold2 (Jumper et al. 2021), RhoFold+ (Shen et al. 2024)\n\nclass InvariantPointAttention(nn.Module):\n    \"\"\"\n    Simplified IPA: attends over single+pair representations with 3D query/key points.\n\n    Three attention components:\n    1. Standard scalar QKV attention (sequence-level)\n    2. Point attention: negative squared distance in 3D between query/key points\n    3. Pair bias from the pair representation\n\n    Output aggregates: scalar values + point values + pair-derived features.\n    \"\"\"\n    def __init__(self, d_single=D_MODEL, d_pair=D_PAIR, n_head=4,\n                 n_query_points=4, n_value_points=4):\n        super().__init__()\n        self.n_head = n_head\n        self.n_qp = n_query_points\n        self.n_vp = n_value_points\n        self.d_head = d_single // n_head\n\n        # Scalar QKV\n        self.q_proj = nn.Linear(d_single, n_head * self.d_head)\n        self.k_proj = nn.Linear(d_single, n_head * self.d_head)\n        self.v_proj = nn.Linear(d_single, n_head * self.d_head)\n\n        # Pair bias -> one logit per head\n        self.pair_proj = nn.Linear(d_pair, n_head)\n\n        # Point QKV in local frames\n        self.q_points = nn.Linear(d_single, n_head * n_query_points * 3)\n        self.k_points = nn.Linear(d_single, n_head * n_query_points * 3)\n        self.v_points = nn.Linear(d_single, n_head * n_value_points * 3)\n\n        # Learnable weight for point attention per head\n        self.head_weights = nn.Parameter(torch.zeros(n_head))\n\n        # Pair values: aggregate pair features weighted by attention\n        self.d_pair_head = d_pair // n_head\n        self.pair_v_proj = nn.Linear(d_pair, n_head * self.d_pair_head)\n\n        # Output projection: scalar + point + pair outputs concatenated\n        out_dim = (n_head * self.d_head\n                   + n_head * n_value_points * 3\n                   + n_head * self.d_pair_head)\n        self.out_proj = nn.Linear(out_dim, d_single)\n\n    def forward(self, s, z, coords, mask=None):\n        \"\"\"\n        s: (B, L, d_single)\n        z: (B, L, L, d_pair)\n        coords: (B, L, 3)\n        mask: (B, L) boolean, True=valid\n        Returns: (B, L, d_single)\n        \"\"\"\n        B, L, _ = s.shape\n        H, dh = self.n_head, self.d_head\n\n        # Scalar QKV\n        q = self.q_proj(s).reshape(B, L, H, dh)\n        k = self.k_proj(s).reshape(B, L, H, dh)\n        v = self.v_proj(s).reshape(B, L, H, dh)\n\n        # Point QKV -> apply translation to get global points\n        qp = self.q_points(s).reshape(B, L, H, self.n_qp, 3)\n        kp = self.k_points(s).reshape(B, L, H, self.n_qp, 3)\n        vp = self.v_points(s).reshape(B, L, H, self.n_vp, 3)\n\n        c = coords.unsqueeze(2).unsqueeze(3)  # (B, L, 1, 1, 3)\n        qp_g = qp + c  # (B, L, H, n_qp, 3)\n        kp_g = kp + c\n        vp_g = vp + c  # (B, L, H, n_vp, 3)\n\n        # --- Attention logits ---\n        # 1) Scalar: (B,H,L,L)\n        attn_scalar = torch.einsum('blhd,bmhd->bhlm', q, k) / math.sqrt(dh)\n\n        # 2) Point: -0.5 * w_h * sum_p ||q_p - k_p||^2\n        qp_r = qp_g.permute(0, 2, 1, 3, 4)  # (B, H, L, n_qp, 3)\n        kp_r = kp_g.permute(0, 2, 1, 3, 4)  # (B, H, L, n_qp, 3)\n        diff = qp_r.unsqueeze(3) - kp_r.unsqueeze(2)  # (B, H, Lq, Lk, n_qp, 3)\n        sq_dist = (diff ** 2).sum(-1).sum(-1)  # (B, H, L, L)\n        w_h = F.softplus(self.head_weights).view(1, H, 1, 1)\n        attn_point = -0.5 * w_h * sq_dist\n\n        # 3) Pair bias: (B,L,L,d_pair) -> (B,L,L,H) -> (B,H,L,L)\n        pair_bias = self.pair_proj(z).permute(0, 3, 1, 2)\n\n        attn = attn_scalar + attn_point + pair_bias\n\n        if mask is not None:\n            key_mask = mask.unsqueeze(1).unsqueeze(2)  # (B, 1, 1, L)\n            attn = attn.masked_fill(~key_mask, float('-inf'))\n\n        attn = F.softmax(attn, dim=-1)\n        attn = torch.nan_to_num(attn, 0.0)\n\n        # --- Aggregate values ---\n        # 1) Scalar values\n        out_s = torch.einsum('bhlm,bmhd->blhd', attn, v).reshape(B, L, H * dh)\n\n        # 2) Point values: (B,H,L,L) x (B,H,L,n_vp,3) -> (B,H,L,n_vp,3)\n        vp_r = vp_g.permute(0, 2, 1, 3, 4)  # (B, H, L, n_vp, 3)\n        out_p = torch.einsum('bhlm,bhmpc->bhlpc', attn, vp_r)  # (B,H,L,n_vp,3)\n        c_local = coords.unsqueeze(1).unsqueeze(3)  # (B, 1, L, 1, 3)\n        out_p = out_p - c_local  # subtract residue position\n        # CRITICAL: permute H before L before flatten\n        out_p = out_p.permute(0, 2, 1, 3, 4).reshape(B, L, H * self.n_vp * 3)\n\n        # 3) Pair values\n        dph = self.d_pair_head\n        pv = self.pair_v_proj(z)  # (B, L, L, H*dph)\n        pv = pv.reshape(B, L, L, H, dph).permute(0, 3, 1, 2, 4)  # (B,H,L,L,dph)\n        out_z = torch.einsum('bhlm,bhmlp->bhlp', attn, pv.permute(0, 1, 3, 2, 4))\n        out_z = out_z.permute(0, 2, 1, 3).reshape(B, L, H * dph)\n\n        out = torch.cat([out_s, out_p, out_z], dim=-1)\n        return self.out_proj(out)\n\n\nclass StructureModule(nn.Module):\n    \"\"\"IPA-based structure module with gradient checkpointing for L*L VRAM savings.\"\"\"\n    def __init__(self, d_single=D_MODEL, d_pair=D_PAIR, n_layers=N_IPA_LAYERS):\n        super().__init__()\n        self.layers = nn.ModuleList()\n        for _ in range(n_layers):\n            self.layers.append(nn.ModuleDict({\n                'ipa': InvariantPointAttention(d_single, d_pair),\n                'norm1': nn.LayerNorm(d_single),\n                'ff': nn.Sequential(\n                    nn.LayerNorm(d_single),\n                    nn.Linear(d_single, d_single * 4),\n                    nn.GELU(),\n                    nn.Dropout(DROPOUT),\n                    nn.Linear(d_single * 4, d_single),\n                ),\n                'norm2': nn.LayerNorm(d_single),\n            }))\n        self.coord_head = nn.Sequential(\n            nn.LayerNorm(d_single),\n            nn.Linear(d_single, d_single // 2),\n            nn.GELU(),\n            nn.Linear(d_single // 2, 3),\n        )\n        nn.init.zeros_(self.coord_head[-1].weight)\n        nn.init.zeros_(self.coord_head[-1].bias)\n\n    def _ipa_step(self, layer, s, z, coords, mask):\n        \"\"\"One IPA layer -- wrapped for gradient checkpointing.\"\"\"\n        s = s + layer['ipa'](layer['norm1'](s), z, coords, mask)\n        s = s + layer['ff'](layer['norm2'](s))\n        delta = self.coord_head(s)\n        coords = coords + delta\n        return s, coords\n\n    def forward(self, s, z, coords_init, mask=None):\n        coords = coords_init\n        for layer in self.layers:\n            if self.training and coords.requires_grad:\n                s, coords = ckpt_fn(\n                    self._ipa_step, layer, s, z, coords, mask,\n                    use_reentrant=False)\n            else:\n                s, coords = self._ipa_step(layer, s, z, coords, mask)\n        return coords.clamp(-COORD_CLAMP, COORD_CLAMP), s\n\nprint(f\"IPA structure module: {N_IPA_LAYERS} layers, gradient checkpointing ON\")","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:56:13.749829Z","iopub.status.busy":"2026-03-25T16:56:13.749417Z","iopub.status.idle":"2026-03-25T16:56:13.768939Z","shell.execute_reply":"2026-03-25T16:56:13.768092Z"},"papermill":{"duration":0.033421,"end_time":"2026-03-25T16:56:13.770305","exception":false,"start_time":"2026-03-25T16:56:13.736884","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"b3aadaf4","cell_type":"code","source":"# Full Model: RNA-FM + Single Track + Pair Track + Structure Module\n\nclass RPBTransformerEncoderLayer(nn.TransformerEncoderLayer):\n    \"\"\"TransformerEncoderLayer with Relative Position Bias.\"\"\"\n    def __init__(self, d_model, nhead, max_rel_dist=64, **kwargs):\n        super().__init__(d_model=d_model, nhead=nhead, **kwargs)\n        self._nhead = nhead\n        self.max_rel_dist = max_rel_dist\n        self.rel_pos_bias = nn.Embedding(2 * max_rel_dist + 1, nhead)\n\n    def _get_rel_bias(self, L, device):\n        pos = torch.arange(L, device=device)\n        rel = (pos[:, None] - pos[None, :]).clamp(-self.max_rel_dist, self.max_rel_dist)\n        rel = rel + self.max_rel_dist\n        return self.rel_pos_bias(rel).permute(2, 0, 1)  # (nhead, L, L)\n\n    def _sa_block(self, x, attn_mask, key_padding_mask, **kwargs):\n        B, L, _ = x.shape\n        rel_bias = self._get_rel_bias(L, x.device)\n        rel_bias = rel_bias.unsqueeze(0).expand(B, -1, -1, -1).reshape(B * self._nhead, L, L)\n        if attn_mask is not None:\n            rel_bias = rel_bias + attn_mask\n        out = self.self_attn(x, x, x, attn_mask=rel_bias,\n                             key_padding_mask=key_padding_mask, need_weights=False)[0]\n        return self.dropout1(out)\n\n\nclass RNAFoldModel(nn.Module):\n    \"\"\"\n    Full RNA 3D folding model:\n    1. RNA-FM embeddings (640-dim, frozen) + handcrafted features (33-dim)\n    2. Single-rep Transformer with RPB (4 layers, d=256)\n    3. Pair-rep track with triangle attention (2 rounds)\n    4. IPA structure module (4 layers) → 3D coordinates\n    5. Recycling: coords fed back for iterative refinement\n    \"\"\"\n    def __init__(self, rna_fm, feature_dim=FEATURE_DIM, d_model=D_MODEL, d_pair=D_PAIR):\n        super().__init__()\n        self.rna_fm = rna_fm  # frozen\n\n        # Project RNA-FM embeddings + handcrafted features → d_model\n        self.fm_proj = nn.Linear(RNA_FM_DIM, d_model)\n        self.feat_proj = nn.Linear(feature_dim, d_model)\n        self.fuse = nn.Sequential(\n            nn.LayerNorm(d_model * 2),\n            nn.Linear(d_model * 2, d_model),\n            nn.GELU(),\n            nn.Dropout(DROPOUT),\n        )\n\n        # Single-representation transformer\n        enc_layer = RPBTransformerEncoderLayer(\n            d_model=d_model, nhead=N_HEAD, dim_feedforward=DIM_FF,\n            dropout=DROPOUT, batch_first=True, norm_first=True)\n        self.single_transformer = nn.TransformerEncoder(enc_layer, num_layers=N_SINGLE_LAYERS)\n\n        # Pair track\n        self.pair_track = PairTrack(d_single=d_model, d_pair=d_pair, n_layers=N_PAIR_LAYERS)\n\n        # Recycle projections\n        self.recycle_s_norm = nn.LayerNorm(d_model)\n        self.recycle_z_norm = nn.LayerNorm(d_pair)\n        self.recycle_coord_norm = nn.LayerNorm(3)\n        self.recycle_s_proj = nn.Linear(d_model + 3, d_model)\n\n        # Structure module\n        self.structure_module = StructureModule(d_single=d_model, d_pair=d_pair,\n                                                n_layers=N_IPA_LAYERS)\n\n        # Distogram head for auxiliary loss\n        self.distogram_head = nn.Sequential(\n            nn.LayerNorm(d_pair),\n            nn.Linear(d_pair, 32),  # 32 distance bins\n        )\n\n    def forward(self, features, sequences, src_key_padding_mask=None, n_recycles=4):\n        \"\"\"\n        features: (B, L, FEATURE_DIM) — handcrafted features\n        sequences: list of str — raw nucleotide sequences for RNA-FM\n        \"\"\"\n        B, L, _ = features.shape\n        mask = ~src_key_padding_mask if src_key_padding_mask is not None else None\n\n        # 1. RNA-FM embeddings (frozen)\n        with torch.no_grad():\n            fm_emb = self.rna_fm(sequences)  # (B, L_fm, 640)\n        # Pad/trim to match L\n        if fm_emb.shape[1] < L:\n            pad = torch.zeros(B, L - fm_emb.shape[1], RNA_FM_DIM, device=fm_emb.device)\n            fm_emb = torch.cat([fm_emb, pad], dim=1)\n        elif fm_emb.shape[1] > L:\n            fm_emb = fm_emb[:, :L, :]\n\n        # Transfer FM embeddings to same device as features (handles multi-GPU)\n        fm_emb = fm_emb.to(features.device)\n\n        # 2. Fuse RNA-FM + handcrafted\n        fm_proj = self.fm_proj(fm_emb)\n        feat_proj = self.feat_proj(features)\n        s_base = self.fuse(torch.cat([fm_proj, feat_proj], dim=-1))\n\n        # 3. Recycling loop\n        coords_prev = None\n        z_prev = None\n        for cycle in range(n_recycles):\n            s = s_base\n            if coords_prev is not None:\n                # Inject previous coords into single rep\n                c = coords_prev.detach()\n                if mask is not None:\n                    valid = mask.float()\n                    n_valid = valid.sum(dim=1, keepdim=True).unsqueeze(-1).clamp(min=1.0)\n                    centroid = (c * valid.unsqueeze(-1)).sum(dim=1, keepdim=True) / n_valid\n                    c = c - centroid\n                else:\n                    c = c - c.mean(dim=1, keepdim=True)\n                c = self.recycle_coord_norm(c)\n                s = self.recycle_s_proj(torch.cat([s, c], dim=-1))\n\n            # Single track\n            pad_mask = src_key_padding_mask\n            s = self.single_transformer(s, src_key_padding_mask=pad_mask)\n\n            # Pair track\n            pair_init = z_prev.detach() if z_prev is not None else None\n            if pair_init is not None:\n                pair_init = self.recycle_z_norm(pair_init)\n            z = self.pair_track(s, pair_init=pair_init, mask=mask)\n\n            # Structure module\n            if coords_prev is None:\n                coords_init = torch.zeros(B, L, 3, device=s.device)\n            else:\n                coords_init = coords_prev.detach()\n            coords_prev, s_out = self.structure_module(s, z, coords_init, mask)\n            z_prev = z\n\n        # Mean-centre final coords\n        if mask is not None:\n            valid = mask.float()\n            n_valid = valid.sum(dim=1, keepdim=True).unsqueeze(-1).clamp(min=1.0)\n            centroid = (coords_prev * valid.unsqueeze(-1)).sum(dim=1, keepdim=True) / n_valid\n            coords_prev = coords_prev - centroid\n        else:\n            coords_prev = coords_prev - coords_prev.mean(dim=1, keepdim=True)\n\n        # Distogram logits\n        dist_logits = self.distogram_head(z)  # (B, L, L, 32)\n\n        return coords_prev, dist_logits\n\n\n# Instantiate\nmodel = RNAFoldModel(rna_fm, feature_dim=FEATURE_DIM).to(DEVICE)\nn_trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\nn_total = sum(p.numel() for p in model.parameters())\nprint(f\"\\u2713 RNAFoldModel  |  Trainable: {n_trainable/1e6:.1f}M  |  Total: {n_total/1e6:.1f}M\")\nprint(f\"  Single: d_model={D_MODEL}, {N_SINGLE_LAYERS} layers, {N_HEAD} heads, ff={DIM_FF}\")\nprint(f\"  Pair:   d_pair={D_PAIR}, {N_PAIR_LAYERS} triangle rounds\")\nprint(f\"  IPA:    {N_IPA_LAYERS} layers, 4 query/value points per head\")\nprint(f\"  RNA-FM: {RNA_FM_DIM}-dim frozen embeddings\")\n\n# --- Multi-GPU support (2xT4) ---\nN_GPUS = torch.cuda.device_count()\nUSE_MULTI_GPU = N_GPUS >= 2\nif USE_MULTI_GPU:\n    model = nn.DataParallel(model, device_ids=list(range(N_GPUS)))\n    print(f\"✓ Multi-GPU: {N_GPUS} GPUs detected - model wrapped with DataParallel\")\nelse:\n    print(f\"  GPUs: {N_GPUS} - single-GPU mode\")\n\n# --- Multi-GPU support (2xT4) ---\nN_GPUS = torch.cuda.device_count()\nUSE_MULTI_GPU = N_GPUS >= 2\nif USE_MULTI_GPU:\n    model = nn.DataParallel(model, device_ids=list(range(N_GPUS)))\n    print(f\"\\u2713 Multi-GPU: {N_GPUS} GPUs detected - model wrapped with DataParallel\")\nelse:\n\n    print(f\"  GPUs: {N_GPUS} \\u2014 single-GPU mode\")    ","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:56:13.794672Z","iopub.status.busy":"2026-03-25T16:56:13.7942Z","iopub.status.idle":"2026-03-25T16:56:13.892605Z","shell.execute_reply":"2026-03-25T16:56:13.891798Z"},"papermill":{"duration":0.1119,"end_time":"2026-03-25T16:56:13.894203","exception":false,"start_time":"2026-03-25T16:56:13.782303","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"255a7e72","cell_type":"code","source":"class ModelEMA:\n    \"\"\"EMA of model weights for stable inference (Polyak averaging).\"\"\"\n    def __init__(self, model, decay=0.999):\n        self.decay = decay\n        self.shadow = {n: p.data.clone() for n, p in model.named_parameters() if p.requires_grad}\n        self.backup = {}\n    def update(self, model):\n        for n, p in model.named_parameters():\n            if p.requires_grad and n in self.shadow:\n                self.shadow[n].mul_(self.decay).add_(p.data, alpha=1 - self.decay)\n    def apply_shadow(self, model):\n        self.backup = {}\n        for n, p in model.named_parameters():\n            if p.requires_grad and n in self.shadow:\n                self.backup[n] = p.data.clone()\n                p.data.copy_(self.shadow[n])\n    def restore(self, model):\n        for n, p in model.named_parameters():\n            if p.requires_grad and n in self.backup:\n                p.data.copy_(self.backup[n])\n        self.backup = {}\n\nema = ModelEMA(model, decay=0.999)\nprint(f\"✓ EMA initialized (decay=0.999, tracking {len(ema.shadow)} parameters)\")","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:56:13.919427Z","iopub.status.busy":"2026-03-25T16:56:13.918697Z","iopub.status.idle":"2026-03-25T16:56:13.934011Z","shell.execute_reply":"2026-03-25T16:56:13.933112Z"},"papermill":{"duration":0.029271,"end_time":"2026-03-25T16:56:13.935398","exception":false,"start_time":"2026-03-25T16:56:13.906127","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"1f659176","cell_type":"code","source":"# Loss Functions: FAPE + dRMSD + Soft TM-score + Distogram CE\n\ndef fape_loss(pred, target, mask, clamp=10.0):\n    \"\"\"Frame Aligned Point Error — clamped per-residue position error after Kabsch.\"\"\"\n    losses = []\n    for i in range(pred.shape[0]):\n        m = mask[i]; p = pred[i][m]; t = target[i][m]\n        if p.shape[0] < 3:\n            continue\n        p_c = p - p.mean(0); t_c = t - t.mean(0)\n        try:\n            H = t_c.T @ p_c\n            U, _, Vh = torch.linalg.svd(H, full_matrices=False)\n            d = torch.det(U @ Vh)\n            flip = torch.cat([torch.ones(2, device=p.device, dtype=p.dtype),\n                              d.sign().unsqueeze(0)])\n            R = (U * flip.unsqueeze(0)) @ Vh\n            errs = torch.sqrt(((p_c @ R.T - t_c) ** 2).sum(-1) + 1e-8)\n            errs = torch.clamp(errs, max=clamp)\n            l = errs.mean()\n            if torch.isfinite(l):\n                losses.append(l)\n        except Exception:\n            pass\n    if not losses:\n        return pred.sum() * 0.0\n    return torch.stack(losses).mean()\n\n\ndef dRMSD_loss(pred, target, mask, max_residues=128, cutoff=12.0):\n    \"\"\"dRMSD over pairs within cutoff (lDDT-inspired).\"\"\"\n    losses = []\n    for i in range(pred.shape[0]):\n        m = mask[i]; p = pred[i][m]; t = target[i][m]; Li = p.shape[0]\n        if Li < 2:\n            continue\n        if Li > max_residues:\n            idx = torch.randperm(Li, device=p.device)[:max_residues]\n            p, t = p[idx], t[idx]\n        pd = torch.cdist(p.unsqueeze(0), p.unsqueeze(0)).squeeze(0)\n        td = torch.cdist(t.unsqueeze(0), t.unsqueeze(0)).squeeze(0)\n        near = td < cutoff\n        diff_sq = ((pd - td) ** 2) * near.float()\n        l = (diff_sq.sum() / near.float().sum().clamp(min=1.0)).sqrt()\n        if torch.isfinite(l):\n            losses.append(l)\n    if not losses:\n        return pred.sum() * 0.0\n    return torch.stack(losses).mean()\n\n\ndef soft_tm_loss(pred, target, mask):\n    \"\"\"Differentiable soft TM-score loss: 1 - mean(1/(1 + d^2/d0^2)).\"\"\"\n    losses = []\n    for i in range(pred.shape[0]):\n        m = mask[i]; p = pred[i][m]; t = target[i][m]; Li = p.shape[0]\n        if Li < 4:\n            continue\n        if   Li < 12: d0 = 0.3\n        elif Li < 16: d0 = 0.4\n        elif Li < 20: d0 = 0.5\n        elif Li < 24: d0 = 0.6\n        elif Li < 30: d0 = 0.7\n        else:         d0 = 1.24 * (Li - 15) ** (1/3) - 1.8\n        d0 = max(d0, 0.5)\n        p_c = p - p.mean(0); t_c = t - t.mean(0)\n        try:\n            H = t_c.T @ p_c\n            U, _, Vh = torch.linalg.svd(H, full_matrices=False)\n            d = torch.det(U @ Vh)\n            flip = torch.cat([torch.ones(2, device=p.device, dtype=p.dtype),\n                              d.sign().unsqueeze(0)])\n            R = (U * flip.unsqueeze(0)) @ Vh\n            p_aligned = p_c @ R.T\n        except Exception:\n            p_aligned = p_c\n        d_sq = ((p_aligned - t_c) ** 2).sum(-1)\n        soft_tm = (1.0 / (1.0 + d_sq / (d0 ** 2))).mean()\n        l = 1.0 - soft_tm\n        if torch.isfinite(l):\n            losses.append(l)\n    if not losses:\n        return pred.sum() * 0.0\n    return torch.stack(losses).mean()\n\n\ndef distogram_loss(dist_logits, target_coords, mask, n_bins=32, v_min=2.0, v_max=22.0):\n    \"\"\"Cross-entropy on binned inter-residue distances — auxiliary pair loss.\"\"\"\n    device = dist_logits.device\n    bin_edges = torch.linspace(v_min, v_max, n_bins + 1, device=device)\n    losses = []\n    B = target_coords.shape[0]\n    for i in range(B):\n        m = mask[i]\n        Li = m.sum().item()\n        if Li < 4:\n            continue\n        t = target_coords[i, :Li]  # (Li, 3) — guaranteed valid since mask[:Li]=True\n        logits = dist_logits[i, :Li, :Li, :]  # (Li, Li, n_bins)\n\n        # Subsample for memory efficiency\n        if Li > 128:\n            idx = torch.randperm(Li, device=device)[:128]\n            t = t[idx]\n            logits = logits[idx][:, idx]\n            Li = 128\n\n        true_dist = torch.cdist(t.unsqueeze(0), t.unsqueeze(0)).squeeze(0)  # (Li, Li)\n        bins = torch.bucketize(true_dist, bin_edges[1:-1]).clamp(0, n_bins - 1)\n\n        l = F.cross_entropy(logits.reshape(-1, n_bins), bins.reshape(-1))\n        if torch.isfinite(l):\n            losses.append(l)\n    if not losses:\n        return dist_logits.sum() * 0.0\n    return torch.stack(losses).mean()\n\n\ndef combined_loss(pred_coords, dist_logits, target, mask, epoch=None):\n    \"\"\"Multi-term loss with TM-score warm-up.\"\"\"\n    TM_WARMUP = 5\n    BLEND_LEN = 5\n    ep = epoch or 999\n\n    l_fape = fape_loss(pred_coords, target, mask)\n    l_drmsd = dRMSD_loss(pred_coords, target, mask)\n    l_disto = distogram_loss(dist_logits, target, mask)\n\n    if ep <= TM_WARMUP:\n        return 0.40 * l_fape + 0.30 * l_drmsd + 0.30 * l_disto\n    elif ep <= TM_WARMUP + BLEND_LEN:\n        alpha = (ep - TM_WARMUP) / BLEND_LEN\n        l_tm = soft_tm_loss(pred_coords, target, mask)\n        w_fape  = 0.40 * (1 - alpha) + 0.30 * alpha\n        w_drmsd = 0.30 * (1 - alpha) + 0.25 * alpha\n        w_disto = 0.30 * (1 - alpha) + 0.20 * alpha\n        w_tm    = 0.25 * alpha\n        return w_fape * l_fape + w_drmsd * l_drmsd + w_disto * l_disto + w_tm * l_tm\n    else:\n        l_tm = soft_tm_loss(pred_coords, target, mask)\n        return 0.30 * l_fape + 0.25 * l_drmsd + 0.20 * l_disto + 0.25 * l_tm\n\nprint(\"✓ Losses: FAPE + dRMSD + soft TM-score + distogram CE\")","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:56:13.960786Z","iopub.status.busy":"2026-03-25T16:56:13.960451Z","iopub.status.idle":"2026-03-25T16:56:13.981431Z","shell.execute_reply":"2026-03-25T16:56:13.980542Z"},"papermill":{"duration":0.035288,"end_time":"2026-03-25T16:56:13.983038","exception":false,"start_time":"2026-03-25T16:56:13.94775","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"7d739d52","cell_type":"markdown","source":"## 4. Training\n\n| Parameter | Value |\n|---|---|\n| Optimizer | AdamW, `lr=3e-4`, `weight_decay=1e-4` |\n| Schedule | Linear warmup (5 epochs) → cosine annealing |\n| Grad accumulation | 8 steps (effective batch = 8) |\n| Grad clipping | `max_norm=1.0` |\n| Recycle curriculum | 1→3 cycles (`min(3, max(1, epoch // 10))`) |\n| Early stopping | Patience = 8 |\n| EMA | decay=0.999 |\n| Mixed precision | AMP with GradScaler |\n| RNA-FM backbone | Frozen (no gradients) |","metadata":{"papermill":{"duration":0.011387,"end_time":"2026-03-25T16:56:14.006053","exception":false,"start_time":"2026-03-25T16:56:13.994666","status":"completed"},"tags":[]}},{"id":"41819149","cell_type":"code","source":"EPOCHS     = 45\nLR         = 3e-4\nPATIENCE   = 8\nGRAD_ACCUM = 8   # effective batch = 8 (1 * 8)\nWARMUP     = 5\nGRAD_CLIP  = 1.0\nCKPT       = OUTPUT_PATH / 'best_model.pt'\n\nif CKPT.exists():\n    CKPT.unlink()\n    print(f\"Removed stale checkpoint: {CKPT}\")\n\n# Only optimize trainable (non-frozen RNA-FM) parameters\ntrainable_params = [p for p in model.parameters() if p.requires_grad]\noptimizer = optim.AdamW(trainable_params, lr=LR, weight_decay=1e-4)\nscaler = GradScaler()\n\n\ndef lr_lambda(ep):\n    if ep < WARMUP:\n        return (ep + 1) / WARMUP\n    prog = (ep - WARMUP) / max(EPOCHS - WARMUP, 1)\n    return 0.5 * (1.0 + math.cos(math.pi * prog))\n\nscheduler = optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)\n\nbest_val_loss, epochs_no_improve = 1e9, 0\ntrain_losses, val_losses = [], []\nlast_val_loss = float('nan')\n\ntrain_loader = DataLoader(train_dataset, batch_sampler=train_sampler, collate_fn=pad_collate, num_workers=2, pin_memory=True)\nval_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, collate_fn=pad_collate, num_workers=2, pin_memory=True)\n\nprint(f\"Training  |  epochs={EPOCHS}  lr={LR}  warmup={WARMUP}  grad_accum={GRAD_ACCUM}\")\nprint(f\"{'Epoch':>6}  {'n_rec':>5}  {'Train Loss':>11}  {'Val Loss':>11}  {'LR':>9}\")\nprint(\"-\" * 55)\n\nfor epoch in range(1, EPOCHS + 1):\n    n_rec = min(3, max(1, epoch // 10))\n\n    model.train()\n    get_underlying_model(model).rna_fm.model.eval()\n    t_loss, n_steps, pending = 0.0, 0, 0\n    optimizer.zero_grad(set_to_none=True)\n\n    for _, (feat, coords, mask, ids, seqs) in enumerate(train_loader):\n        feat, coords, mask = feat.to(DEVICE), coords.to(DEVICE), mask.to(DEVICE)\n        pad_mask = ~mask\n        sequences = list(seqs)\n\n        with autocast(device_type='cuda' if DEVICE == 'cuda' else 'cpu'):\n            pred_coords, dist_logits = model(feat, sequences,\n                                             src_key_padding_mask=pad_mask,\n                                             n_recycles=n_rec)\n            loss = combined_loss(pred_coords, dist_logits, coords, mask, epoch=epoch) / GRAD_ACCUM\n\n        if not torch.isfinite(loss):\n            continue\n\n        scaler.scale(loss).backward()\n        pending += 1\n\n        if pending >= GRAD_ACCUM:\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(trainable_params, GRAD_CLIP)\n            scaler.step(optimizer)\n            scaler.update()\n            ema.update(model)\n            optimizer.zero_grad(set_to_none=True)\n            pending = 0\n\n        t_loss += loss.item() * GRAD_ACCUM\n        n_steps += 1\n\n    if pending > 0:\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(trainable_params, GRAD_CLIP)\n        scaler.step(optimizer)\n        scaler.update()\n        ema.update(model)\n        optimizer.zero_grad(set_to_none=True)\n\n    t_loss /= max(n_steps, 1)\n\n    # Validation — run every 3 epochs to save time\n    do_val = (epoch == 1) or (epoch % 3 == 0)\n    if do_val:\n        val_rec = max(n_rec, 2)\n        ema.apply_shadow(model)\n        model.eval()\n        v_loss, v_count = 0.0, 0\n        with torch.no_grad():\n            for feat, coords, mask, ids, seqs in val_loader:\n                feat, coords, mask = feat.to(DEVICE), coords.to(DEVICE), mask.to(DEVICE)\n                pad_mask = ~mask\n                sequences = list(seqs)\n                with autocast(device_type='cuda' if DEVICE == 'cuda' else 'cpu'):\n                    pred_coords, dist_logits = model(feat, sequences,\n                                                     src_key_padding_mask=pad_mask,\n                                                     n_recycles=val_rec)\n                    batch_v = combined_loss(pred_coords, dist_logits, coords, mask, epoch=epoch)\n                if torch.isfinite(batch_v):\n                    v_loss += batch_v.item()\n                    v_count += 1\n        v_loss = v_loss / max(v_count, 1) if v_count > 0 else float('inf')\n        last_val_loss = v_loss\n        ema.restore(model)\n    else:\n        v_loss = float('nan')\n\n    scheduler.step()\n    train_losses.append(t_loss)\n    val_losses.append(v_loss)\n    lr_now = optimizer.param_groups[0]['lr']\n\n    display_val_loss = v_loss if math.isfinite(v_loss) else last_val_loss\n    v_str = f\"{display_val_loss:>11.4f}\" if math.isfinite(display_val_loss) else \"          -\"\n    print(f\"{epoch:>6}  {n_rec:>5}  {t_loss:>11.4f}  {v_str}  {lr_now:>9.2e}\")\n\n    if do_val and math.isfinite(v_loss):\n        if v_loss < best_val_loss:\n            best_val_loss = v_loss\n            epochs_no_improve = 0\n            ema.apply_shadow(model)\n            torch.save(model.state_dict(), CKPT)\n            ema.restore(model)\n        else:\n            epochs_no_improve += 1\n            if epochs_no_improve >= PATIENCE:\n                print(f\"\\nEarly stopping at epoch {epoch}  (best val loss: {best_val_loss:.4f})\")\n                break\n\nif not CKPT.exists():\n    print(\"No best checkpoint saved -- saving final weights.\")\n    torch.save(model.state_dict(), CKPT)\n\nmodel.load_state_dict(torch.load(CKPT, map_location=DEVICE, weights_only=True))\nema = ModelEMA(model, decay=0.0)\nprint(f\"\\nBest val combined loss: {best_val_loss:.6f}\")\n\ngc.collect()\ntorch.cuda.empty_cache()\n\nplt.figure(figsize=(9, 3))\nepochs_range = list(range(1, len(train_losses) + 1))\nfin_t = [(e, v) for e, v in zip(epochs_range, train_losses) if math.isfinite(v)]\nfin_v = [(e, v) for e, v in zip(epochs_range, val_losses) if math.isfinite(v)]\nif fin_t: plt.plot(*zip(*fin_t), label='Train')\nif fin_v: plt.plot(*zip(*fin_v), label='Val')\nplt.xlabel('Epoch'); plt.ylabel('Combined Loss')\nplt.title('Training Curve'); plt.legend(); plt.tight_layout(); plt.show()","metadata":{"execution":{"iopub.execute_input":"2026-03-25T16:56:14.030302Z","iopub.status.busy":"2026-03-25T16:56:14.030041Z","iopub.status.idle":"2026-03-26T00:06:37.280811Z","shell.execute_reply":"2026-03-26T00:06:37.279904Z"},"papermill":{"duration":25823.265659,"end_time":"2026-03-26T00:06:37.282981","exception":false,"start_time":"2026-03-25T16:56:14.017322","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"43c1ff51","cell_type":"code","source":"# Full-data fine-tuning on train + validation combined\nFULL_FINETUNE_EPOCHS = 2\nFULL_FINETUNE_LR = 1e-4\n\nif FULL_FINETUNE_EPOCHS > 0:\n    print(f\"\\nStarting full-data fine-tune: epochs={FULL_FINETUNE_EPOCHS}, lr={FULL_FINETUNE_LR}\")\n\n    full_seq_df = pd.concat([train_seq, val_seq], ignore_index=True)\n    full_seq_df = full_seq_df.drop_duplicates(subset=[SEQ_ID_COL]).reset_index(drop=True)\n    full_labels = pd.concat([train_labels, val_labels], ignore_index=True)\n\n    full_dataset = RNADataset(full_seq_df, full_labels, max_len=MAX_TRAIN_LEN, name='full', augment=True)\n    full_sampler = BucketBatchSampler(full_dataset, batch_size=BATCH_SIZE, bucket_size=100)\n    full_loader = DataLoader(full_dataset, batch_sampler=full_sampler, collate_fn=pad_collate, num_workers=2, pin_memory=True)\n\n    ema = ModelEMA(model, decay=0.999)\n    ft_trainable = [p for p in model.parameters() if p.requires_grad]\n    ft_opt = optim.AdamW(ft_trainable, lr=FULL_FINETUNE_LR, weight_decay=1e-4)\n\n    model.train()\n    get_underlying_model(model).rna_fm.model.eval()\n    for ep in range(1, FULL_FINETUNE_EPOCHS + 1):\n        n_rec = 3\n        running, n_steps, pending = 0.0, 0, 0\n        ft_opt.zero_grad(set_to_none=True)\n        for feat, coords, mask, ids, seqs in full_loader:\n            feat, coords, mask = feat.to(DEVICE), coords.to(DEVICE), mask.to(DEVICE)\n            pad_mask = ~mask\n            sequences = list(seqs)\n            with autocast(device_type='cuda' if DEVICE == 'cuda' else 'cpu'):\n                pred_coords, dist_logits = model(feat, sequences,\n                                                 src_key_padding_mask=pad_mask,\n                                                 n_recycles=n_rec)\n                loss = combined_loss(pred_coords, dist_logits, coords, mask, epoch=999) / GRAD_ACCUM\n            if not torch.isfinite(loss):\n                continue\n            scaler.scale(loss).backward()\n            pending += 1\n            if pending >= GRAD_ACCUM:\n                scaler.unscale_(ft_opt)\n                torch.nn.utils.clip_grad_norm_(ft_trainable, GRAD_CLIP)\n                scaler.step(ft_opt)\n                scaler.update()\n                ema.update(model)\n                ft_opt.zero_grad(set_to_none=True)\n                pending = 0\n            running += loss.item() * GRAD_ACCUM\n            n_steps += 1\n\n        if pending > 0:\n            scaler.unscale_(ft_opt)\n            torch.nn.utils.clip_grad_norm_(ft_trainable, GRAD_CLIP)\n            scaler.step(ft_opt)\n            scaler.update()\n            ema.update(model)\n            ft_opt.zero_grad(set_to_none=True)\n\n        ep_loss = running / max(n_steps, 1)\n        print(f\"  [full-ft] epoch {ep:02d}/{FULL_FINETUNE_EPOCHS}  loss={ep_loss:.4f}  n_rec={n_rec}\")\n\n    ema.apply_shadow(model)\n    torch.save(model.state_dict(), CKPT)\n    ema.restore(model)\n    # Load EMA checkpoint so inference uses best weights\n    model.load_state_dict(torch.load(CKPT, map_location=DEVICE, weights_only=True))\n    print(f\"Full-data fine-tune complete. Best EMA weights loaded.\")\n\n    # Free datasets no longer needed\n    del full_dataset, full_loader, full_sampler\n    del train_dataset, train_loader, train_sampler\n    gc.collect()\n    torch.cuda.empty_cache()\n    print(\"  Memory freed: deleted all datasets and loaders\")\n","metadata":{"execution":{"iopub.execute_input":"2026-03-26T00:06:37.313237Z","iopub.status.busy":"2026-03-26T00:06:37.313003Z","iopub.status.idle":"2026-03-26T00:32:22.747459Z","shell.execute_reply":"2026-03-26T00:32:22.74639Z"},"papermill":{"duration":1545.451038,"end_time":"2026-03-26T00:32:22.74914","exception":false,"start_time":"2026-03-26T00:06:37.298102","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"8692d7af","cell_type":"markdown","source":"## 5. Validation\n\nEvaluate TM-score on the held-out validation set using the best checkpoint.\n\n| TM-score | Meaning |\n|---|---|\n| < 0.2 | Random |\n| 0.2–0.5 | Weak similarity |\n| > 0.5 | Same fold |\n| > 0.9 | Near-identical |","metadata":{"papermill":{"duration":0.014277,"end_time":"2026-03-26T00:32:22.778765","exception":false,"start_time":"2026-03-26T00:32:22.764488","status":"completed"},"tags":[]}},{"id":"e9a2501d","cell_type":"code","source":"# Evaluate TM-score on validation set\n# Defensive guard: rebuild validation loader if this cell is run out of order\nif 'val_loader' not in globals():\n    if 'val_dataset' not in globals():\n        val_dataset = RNADataset(val_seq, val_labels, max_len=MAX_TRAIN_LEN, name='val', augment=False)\n    val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, collate_fn=pad_collate, num_workers=2, pin_memory=True)\n\nmodel.eval()\ntm_scores = []\nfor feat, coord_true_t, mask, ids, seqs in val_loader:\n    feat = feat.to(DEVICE)\n    pad_mask = (~mask).to(DEVICE)\n    sequences = list(seqs)\n    with torch.no_grad():\n        with autocast(device_type='cuda' if DEVICE == 'cuda' else 'cpu'):\n            preds, _ = model(feat, sequences, src_key_padding_mask=pad_mask)\n        preds = preds.cpu().numpy()\n    masks_np = mask.numpy()\n    coord_true_np = coord_true_t.numpy()\n    for i in range(len(ids)):\n        m = masks_np[i]\n        p = preds[i][m]\n        ref = coord_true_np[i][m]\n        tm_scores.append(tm_score(p, ref))\n\nprint(f\"Validation TM-score:  \"\n      f\"mean={np.mean(tm_scores):.4f}  \"\n      f\"median={np.median(tm_scores):.4f}  \"\n      f\"min={np.min(tm_scores):.4f}  \"\n      f\"max={np.max(tm_scores):.4f}\")\n\nplt.figure(figsize=(8, 3))\nplt.hist(tm_scores, bins=20, edgecolor='k')\nplt.axvline(np.mean(tm_scores), color='red', ls='--',\n            label=f'Mean={np.mean(tm_scores):.3f}')\nplt.xlabel('TM-score'); plt.ylabel('Count')\nplt.title('Validation TM-score Distribution'); plt.legend()\nplt.tight_layout(); plt.show()\n","metadata":{"execution":{"iopub.execute_input":"2026-03-26T00:32:22.811076Z","iopub.status.busy":"2026-03-26T00:32:22.810673Z","iopub.status.idle":"2026-03-26T00:32:28.374345Z","shell.execute_reply":"2026-03-26T00:32:28.373462Z"},"papermill":{"duration":5.583542,"end_time":"2026-03-26T00:32:28.376053","exception":false,"start_time":"2026-03-26T00:32:22.792511","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"352431d6","cell_type":"markdown","source":"## 6. Generate 5 Diverse Predictions\n\nBest-of-5 scoring rewards diversity. Strategy:\n\n| Slot | Method | Source |\n|---|---|---|\n| 1 | Deterministic (eval mode, 3 recycles) | Full MSA |\n| 2 | MC Dropout (train mode, 3 recycles) | Full MSA |\n| 3 | Retrieval candidate | Sequence similarity bank |\n| 4 | Template candidate | CIF structural data |\n| 5 | MC Dropout + MSA subsample | Random MSA subset |","metadata":{"papermill":{"duration":0.017719,"end_time":"2026-03-26T00:32:28.41006","exception":false,"start_time":"2026-03-26T00:32:28.392341","status":"completed"},"tags":[]}},{"id":"c797ed3d","cell_type":"code","source":"def predict_5_diverse(model, seq_id, seq, stoichiometry=None, all_sequences=None,\n                      device=DEVICE, model_gpu1=None):\n    \"\"\"5 predictions combining neural, retrieval, and template candidates.\"\"\"\n    mpath = resolve_msa_key(seq_id)\n    L = len(seq)\n    chain_lens = get_chain_lengths(stoichiometry, all_sequences, L)\n\n    # Pre-compute MI features once (O(L²)) — reused for all predictions\n    _mi_cached = msa_mi_features(mpath, L) if (mpath and USE_MI_FEAT) else np.zeros((L, 4), np.float32)\n\n    def build_feat_with_msa_subset(msa_path, max_seqs=None, seed=None):\n        parts = [one_hot(seq), positional_encoding_with_chains(chain_lens, d=16)]\n        if USE_CHAIN_BOUNDARY:\n            parts.append(chain_boundary_feature(chain_lens))\n        if USE_MSA:\n            if msa_path and max_seqs is not None:\n                all_msa_seqs = parse_msa(msa_path)\n                if seed is not None:\n                    rng = random.Random(seed)\n                    sub = rng.sample(all_msa_seqs, min(max_seqs, len(all_msa_seqs)))\n                else:\n                    sub = all_msa_seqs[:max_seqs]\n                freq = np.zeros((L, 4), dtype=np.float32)\n                cnt = np.zeros(L, dtype=np.float32)\n                for s in sub:\n                    ungapped = 0\n                    for ch in s:\n                        if ch in ('-', '.'):\n                            continue\n                        if ungapped >= L:\n                            break\n                        idx = NT_TO_IDX.get(ch.upper(), -1)\n                        if idx >= 0:\n                            freq[ungapped, idx] += 1\n                            cnt[ungapped] += 1\n                        ungapped += 1\n                cnt = np.maximum(cnt, 1)[:, None]\n                parts.append((freq / cnt).astype(np.float32))\n            else:\n                parts.append(msa_features(msa_path, L) if msa_path else np.zeros((L, 4), np.float32))\n        if USE_TEMPLATE:\n            parts.append(template_features(seq_id, L))\n        if USE_MI_FEAT:\n            parts.append(_mi_cached)\n        return np.concatenate(parts, axis=1).astype(np.float32)\n\n    def calibrate_backbone_scale(coords, target_step=6.0):\n        c = coords.astype(np.float32, copy=True)\n        c -= c.mean(axis=0, keepdims=True)\n        if len(c) < 2:\n            return c\n        step = np.linalg.norm(np.diff(c, axis=0), axis=1)\n        step = step[np.isfinite(step)]\n        if step.size == 0:\n            return c\n        mean_step = float(step.mean())\n        if mean_step < 1e-6:\n            return c\n        if mean_step < 4.5 or mean_step > 8.5:\n            scale = np.clip(target_step / mean_step, 0.6, 2.5)\n            c *= scale\n        return c\n\n    def _segments_from_chain_lengths(lengths, total_len):\n        if not lengths:\n            return [(0, total_len)]\n        segments = []\n        pos = 0\n        for ln in lengths:\n            try:\n                seg_len = int(ln)\n            except Exception:\n                continue\n            if seg_len <= 0:\n                continue\n            end = min(pos + seg_len, total_len)\n            segments.append((pos, end))\n            pos = end\n            if pos >= total_len:\n                break\n        if not segments:\n            return [(0, total_len)]\n        if segments[-1][1] < total_len:\n            segments.append((segments[-1][1], total_len))\n        return segments\n\n    def refine_chain_geometry(coords, lengths, passes=3):\n        X = np.asarray(coords, dtype=np.float64).copy()\n        segments = _segments_from_chain_lengths(lengths, len(X))\n        if len(X) < 2:\n            return X.astype(np.float32)\n\n        for _ in range(passes):\n            for s, e in segments:\n                C = X[s:e]\n                L = e - s\n                if L < 3:\n                    continue\n\n                # Keep adjacent residue spacing near an A-form-like backbone step.\n                d = C[1:] - C[:-1]\n                dist = np.linalg.norm(d, axis=1, keepdims=True).clip(1e-6)\n                adj = d * ((6.0 - dist) / dist) * 0.18\n                C[:-1] -= adj\n                C[1:] += adj\n\n                # Gentle second-neighbor correction.\n                d2 = C[2:] - C[:-2]\n                d2n = np.linalg.norm(d2, axis=1, keepdims=True).clip(1e-6)\n                adj2 = d2 * ((10.2 - d2n) / d2n) * 0.08\n                C[:-2] -= adj2\n                C[2:] += adj2\n\n                # Smooth local kinks.\n                C[1:-1] += 0.06 * (0.5 * (C[:-2] + C[2:]) - C[1:-1])\n\n                # Sampled self-avoidance, adapted from the winner notebook.\n                if L >= 20:\n                    idx = (np.linspace(0, L - 1, min(L, 200)).astype(int)\n                           if L > 240 else np.arange(L))\n                    P = C[idx]\n                    diff = P[:, None, :] - P[None, :, :]\n                    dm = np.linalg.norm(diff, axis=2).clip(1e-6)\n                    sep = np.abs(idx[:, None] - idx[None, :])\n                    mask = (sep > 3) & (dm < 4.0)\n                    if np.any(mask):\n                        push = (diff * ((4.0 - dm) / dm)[:, :, None]\n                                * mask[:, :, None]).sum(axis=1)\n                        C[idx] += 0.015 * push\n\n                X[s:e] = C\n\n            X -= X.mean(axis=0, keepdims=True)\n\n        step = np.linalg.norm(np.diff(X, axis=0), axis=1)\n        step = step[np.isfinite(step)]\n        if step.size > 0:\n            mean_step = float(step.mean())\n            if mean_step >= 1e-6 and (mean_step < 4.5 or mean_step > 8.5):\n                X *= np.clip(6.0 / mean_step, 0.6, 2.5)\n        return X.astype(np.float32)\n\n    def candidate_quality(coords):\n        arr = np.asarray(coords, dtype=np.float64)\n        if arr.ndim != 2 or arr.shape[0] < 2:\n            return float('inf')\n        step = np.linalg.norm(np.diff(arr, axis=0), axis=1)\n        step = step[np.isfinite(step)]\n        if step.size == 0:\n            return float('inf')\n        return float(abs(step.mean() - 6.0) + step.std())\n\n    def consensus_candidate(pred_list):\n        if len(pred_list) < 2:\n            return None\n        stack = np.stack(pred_list, axis=0).astype(np.float64)\n        weights = np.array([1.0 / max(candidate_quality(p), 0.1) for p in pred_list], dtype=np.float64)\n        weights = weights / weights.sum()\n        blended = np.tensordot(weights, stack, axes=(0, 0))\n        return refine_chain_geometry(blended, chain_lens, passes=2)\n\n    def _run_model_once(feat_t, subseq, n_rec, use_gpu1=False):\n        \"\"\"Run model on a single (1, L_win, F) tensor. Returns (L_win, 3) numpy.\n        If use_gpu1=True and model_gpu1 exists, runs on cuda:1.\"\"\"\n        with torch.no_grad():\n            if use_gpu1 and model_gpu1 is not None:\n                _f = feat_t.to('cuda:1')\n                with autocast(device_type='cuda'):\n                    pred_coords, _ = model_gpu1(_f, [subseq], n_recycles=n_rec)\n            else:\n                with autocast(device_type='cuda' if DEVICE == 'cuda' else 'cpu'):\n                    pred_coords, _ = model(feat_t, [subseq], n_recycles=n_rec)\n            return pred_coords.float().squeeze(0).cpu().numpy()\n\n    def neural_predict(feat_np, train_mode=False, n_rec=3, gpu1=False):\n        _use_gpu1 = gpu1 and model_gpu1 is not None\n        if _use_gpu1:\n            if train_mode:\n                model_gpu1.train()\n                get_underlying_model(model_gpu1).rna_fm.model.eval()\n            else:\n                model_gpu1.eval()\n        else:\n            if train_mode:\n                model.train()\n                get_underlying_model(model).rna_fm.model.eval()\n            else:\n                model.eval()\n\n        L_seq = feat_np.shape[0]\n        # Short sequences: run directly\n        if L_seq <= INFER_WIN_LEN:\n            x = torch.tensor(feat_np).unsqueeze(0).to(device if not _use_gpu1 else 'cuda:1')\n            return _run_model_once(x, seq, n_rec, use_gpu1=_use_gpu1)\n\n        # Long sequences: sliding window with overlap blending\n        # Each window predicts in its own coordinate frame, so we must\n        # Kabsch-align overlapping regions before blending.\n        win, stride = INFER_WIN_LEN, INFER_STRIDE\n\n        starts = list(range(0, L_seq - win + 1, stride))\n        if not starts or starts[-1] + win < L_seq:\n            starts.append(max(L_seq - win, 0))\n\n        # Predict all windows\n        win_coords = []\n        for s in starts:\n            e = s + win\n            feat_win = torch.tensor(feat_np[s:e]).unsqueeze(0).to(device if not _use_gpu1 else 'cuda:1')\n            subseq = seq[s:e]\n            c = _run_model_once(feat_win, subseq, n_rec, use_gpu1=_use_gpu1)  # (win, 3)\n            win_coords.append((s, e, c))\n\n        # Stitch: align each successive window to the running global coords\n        global_coords = np.full((L_seq, 3), np.nan, dtype=np.float64)\n        global_weight = np.zeros((L_seq, 1), dtype=np.float64)\n\n        # First window anchors the global frame\n        s0, e0, c0 = win_coords[0]\n        w0 = np.minimum(\n            np.arange(1, e0 - s0 + 1, dtype=np.float64),\n            np.arange(e0 - s0, 0, -1, dtype=np.float64),\n        ).reshape(-1, 1)\n        global_coords[s0:e0] = c0\n        global_weight[s0:e0] = w0\n\n        for s, e, c in win_coords[1:]:\n            # Overlap region with existing global coords\n            overlap_mask = ~np.isnan(global_coords[s:e, 0])\n            n_overlap = overlap_mask.sum()\n            if n_overlap >= 3:\n                # Kabsch-align this window's coords to global using the overlap\n                ref = global_coords[s:e][overlap_mask].copy()\n                mov = c[overlap_mask].copy()\n                ref_c = ref - ref.mean(0)\n                mov_c = mov - mov.mean(0)\n                H = mov_c.T @ ref_c\n                U, _, Vt = np.linalg.svd(H)\n                d = np.linalg.det(Vt.T @ U.T)\n                D = np.diag([1.0, 1.0, d])\n                R = (Vt.T @ D @ U.T)\n                c_aligned = (c - mov.mean(0)) @ R.T + ref.mean(0)\n            else:\n                # Not enough overlap — just translate to match\n                if n_overlap > 0:\n                    ref_mean = global_coords[s:e][overlap_mask].mean(0)\n                    mov_mean = c[overlap_mask].mean(0)\n                    c_aligned = c - mov_mean + ref_mean\n                else:\n                    c_aligned = c\n\n            w = np.minimum(\n                np.arange(1, e - s + 1, dtype=np.float64),\n                np.arange(e - s, 0, -1, dtype=np.float64),\n            ).reshape(-1, 1)\n\n            # Weighted blend into global coords\n            for idx_local, idx_global in enumerate(range(s, e)):\n                wt = w[idx_local, 0]\n                if np.isnan(global_coords[idx_global, 0]):\n                    global_coords[idx_global] = c_aligned[idx_local]\n                    global_weight[idx_global] = wt\n                else:\n                    old_wt = global_weight[idx_global, 0]\n                    total = old_wt + wt\n                    global_coords[idx_global] = (\n                        global_coords[idx_global] * old_wt + c_aligned[idx_local] * wt\n                    ) / total\n                    global_weight[idx_global, 0] = total\n\n        # Fill any remaining NaN (shouldn't happen if starts cover [0, L_seq])\n        nan_mask = np.isnan(global_coords[:, 0])\n        if nan_mask.any():\n            idx = np.arange(L_seq, dtype=np.float32)\n            for dim in range(3):\n                series = global_coords[:, dim]\n                finite = np.isfinite(series)\n                if finite.sum() >= 2:\n                    series[~finite] = np.interp(idx[~finite], idx[finite], series[finite])\n                elif finite.sum() == 1:\n                    series[~finite] = series[finite][0]\n                else:\n                    return None\n                global_coords[:, dim] = series\n        return global_coords.astype(np.float32)\n\n    n_msa = len(parse_msa(mpath)) if mpath else 0\n    preds, labels = [], []\n\n    def _repair_candidate(coords):\n        arr = np.asarray(coords, dtype=np.float32)\n        if arr.ndim != 2 or arr.shape[1] != 3:\n            return None\n\n        idx = np.arange(len(arr), dtype=np.float32)\n        repaired = arr.copy()\n        for dim in range(3):\n            series = repaired[:, dim]\n            mask = np.isfinite(series)\n            if mask.sum() >= 2:\n                series[~mask] = np.interp(idx[~mask], idx[mask], series[mask])\n            elif mask.sum() == 1:\n                series[:] = series[mask][0]\n            else:\n                return None\n            repaired[:, dim] = series\n\n        repaired -= repaired.mean(axis=0, keepdims=True)\n        if not np.isfinite(repaired).all():\n            return None\n\n        step = np.linalg.norm(np.diff(repaired, axis=0), axis=1)\n        step = step[np.isfinite(step)]\n        if step.size > 0:\n            mean_step = float(step.mean())\n            if mean_step >= 1e-6 and (mean_step < 4.5 or mean_step > 8.5):\n                repaired *= np.clip(6.0 / mean_step, 0.6, 2.5)\n        return repaired.astype(np.float32)\n\n    def _finalize_candidate(coords):\n        if coords is None:\n            return None\n        candidate = _repair_candidate(calibrate_backbone_scale(coords))\n        if candidate is None or not np.isfinite(candidate).all():\n            return None\n        candidate = refine_chain_geometry(candidate, chain_lens, passes=3)\n        if not np.isfinite(candidate).all():\n            return None\n        return candidate.astype(np.float32)\n\n    def _add_candidate(coords, label):\n        candidate = _finalize_candidate(coords)\n        if candidate is not None:\n            preds.append(candidate)\n            labels.append(label)\n\n    # 1) Deterministic\n    feat_np = build_features(seq_id, seq, stoichiometry=stoichiometry, all_sequences=all_sequences)\n    _add_candidate(neural_predict(feat_np, train_mode=False), 'det(full)')\n\n    # 2) MC-dropout\n    _add_candidate(neural_predict(feat_np, train_mode=True), 'MC(full)')\n\n    # 3) Retrieval\n    r_pred = retrieval_candidate(seq, top_k=3, min_score=0.35)\n    if r_pred is not None:\n        _add_candidate(r_pred, 'retrieval')\n\n    # 4) Template\n    t_pred = template_candidate(seq_id, L, min_coverage_frac=0.20)\n    if t_pred is not None:\n        _add_candidate(t_pred, 'template')\n\n    # Extra actual model passes for 5-prediction coverage (no synthetic jitter)\n    extra_cfgs = [\n        (None, 42, 'MC(extra,42)'),\n        (None, 123, 'MC(extra,123)'),\n        (None, 99, 'MC(extra,99)'),\n        (None, 7, 'MC(extra,7)'),\n        (None, 17, 'MC(extra,17)'),\n    ]\n    if mpath:\n        half = n_msa // 2 if n_msa >= 16 else None\n        qtr = n_msa // 4 if n_msa >= 16 else None\n        extra_cfgs = [\n            (half, 42, 'MC(half,42)'),\n            (qtr, 123, 'MC(qtr,123)'),\n            (half, 99, 'MC(half,99)'),\n            (qtr, 7, 'MC(qtr,7)'),\n            (half, 17, 'MC(half,17)'),\n        ]\n\n    for max_seqs, seed, lb in extra_cfgs:\n        if len(preds) >= 5:\n            break\n        f = build_feat_with_msa_subset(mpath, max_seqs=max_seqs, seed=seed)\n        cand = neural_predict(f, train_mode=True, gpu1=(model_gpu1 is not None))\n        _add_candidate(cand, lb)\n\n    if len(preds) < 5:\n        raise ValueError(f'Unable to assemble 5 valid prediction candidates for {seq_id}; got {len(preds)}')\n\n    model.eval()\n    return preds[:5], labels[:5]\n","metadata":{"execution":{"iopub.execute_input":"2026-03-26T00:32:28.446388Z","iopub.status.busy":"2026-03-26T00:32:28.446095Z","iopub.status.idle":"2026-03-26T00:32:28.501259Z","shell.execute_reply":"2026-03-26T00:32:28.500524Z"},"papermill":{"duration":0.076062,"end_time":"2026-03-26T00:32:28.502794","exception":false,"start_time":"2026-03-26T00:32:28.426732","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"c213ad22","cell_type":"code","source":"# Optional Protenix-backed predictions (requires Protenix package + code dir)\ndef build_input_json(df, json_path):\n    data = []\n    for _, row in df.iterrows():\n        data.append({\n            \"name\": str(row[SEQ_ID_COL]),\n            \"covalent_bonds\": [],\n            \"sequences\": [{\"rnaSequence\": {\"sequence\": str(row[SEQ_COL]), \"count\": 1}}],\n        })\n    with open(json_path, \"w\", encoding=\"utf-8\") as f:\n        json.dump(data, f)\n\ndef build_configs(input_json_path, dump_dir, model_name):\n    from configs.configs_base import configs as configs_base\n    from configs.configs_data import data_configs\n    from configs.configs_inference import inference_configs\n    from configs.configs_model_type import model_configs\n    from protenix.config.config import parse_configs\n\n    base = {**configs_base, **{\"data\": data_configs}, **inference_configs}\n\n    def deep_update(target, patch):\n        for key, value in patch.items():\n            if isinstance(value, dict) and key in target and isinstance(target[key], dict):\n                deep_update(target[key], value)\n            else:\n                target[key] = value\n\n    deep_update(base, model_configs[model_name])\n    arg_str = \" \".join([\n        f\"--model_name {model_name}\",\n        f\"--input_json_path {input_json_path}\",\n        f\"--dump_dir {dump_dir}\",\n        f\"--use_msa {USE_MSA}\",\n        f\"--use_template {USE_TEMPLATE}\",\n        f\"--use_rna_msa {USE_RNA_MSA}\",\n        f\"--sample_diffusion.N_sample {N_SAMPLE}\",\n        f\"--seeds {42}\",\n    ])\n    return parse_configs(configs=base, arg_str=arg_str, fill_required_with_null=True)\n\ndef extract_c1_coords(raw_coords, data, full_seq_len):\n    feat = data[\"input_feature_dict\"]\n    device = raw_coords.device\n    n_sample = raw_coords.shape[0]\n    n_atoms = raw_coords.shape[1]\n    mask = None\n    for key in (\"centre_atom_mask\", \"center_atom_mask\"):\n        if key in feat:\n            m = (feat[key] == 1)\n            if m.sum().item() > 0:\n                mask = m.to(device)\n                break\n    if mask is None and \"atom_to_tokatom_idx\" in feat:\n        idx_tensor = feat[\"atom_to_tokatom_idx\"].long()\n        best_idx, best_diff = 11, float(\"inf\")\n        for candidate in range(20):\n            cnt = (idx_tensor == candidate).sum().item()\n            diff = abs(cnt - full_seq_len)\n            if diff < best_diff:\n                best_diff = diff\n                best_idx = candidate\n        mask = (idx_tensor == best_idx).to(device)\n    if mask is None or mask.sum().item() == 0:\n        stride = max(1, n_atoms // max(full_seq_len, 1))\n        indices = torch.arange(0, n_atoms, stride, device=device)[:full_seq_len]\n        mask = torch.zeros(n_atoms, dtype=torch.bool, device=device)\n        mask[indices] = True\n    coords = raw_coords[:, mask, :].detach().cpu().numpy()\n    if coords.shape[1] == full_seq_len:\n        return coords\n    out = np.zeros((n_sample, full_seq_len, 3), dtype=np.float32)\n    use = min(coords.shape[1], full_seq_len)\n    if use > 0:\n        out[:, :use] = coords[:, :use]\n    if np.all(np.abs(out) < 1e-5):\n        helix = np.zeros((full_seq_len, 3), dtype=np.float32)\n        for i in range(full_seq_len):\n            ang = i * 0.5890\n            helix[i] = [9.4 * np.cos(ang), 9.4 * np.sin(ang), i * 2.81]\n        out = np.tile(helix[None], (n_sample, 1, 1))\n    return out\n\ndef protenix_predictions_for_test(test_df):\n    if not USE_PROTENIX or PROTENIX_CODE_DIR is None or PROTENIX_ROOT_DIR is None:\n        return {}\n\n    try:\n        import sys\n        sys.path.insert(0, str(PROTENIX_CODE_DIR))\n        os.environ[\"PROTENIX_ROOT_DIR\"] = str(PROTENIX_ROOT_DIR)\n        from protenix.data.inference.infer_dataloader import InferenceDataset\n        from runner.inference import (InferenceRunner, update_gpu_compatible_configs, update_inference_configs)\n    except Exception as exc:\n        print(f\"Protenix unavailable, using fallback predictor instead: {exc}\")\n        return {}\n\n    work_dir = OUTPUT_PATH / 'protenix_work'\n    work_dir.mkdir(parents=True, exist_ok=True)\n    input_json_path = str(work_dir / 'protenix_input.json')\n    build_input_json(test_df[[SEQ_ID_COL, SEQ_COL]].copy(), input_json_path)\n    configs = build_configs(input_json_path, str(work_dir / 'outputs'), MODEL_NAME)\n    configs = update_gpu_compatible_configs(configs)\n    runner = InferenceRunner(configs)\n    dataset = InferenceDataset(configs)\n    seq_by_id = dict(zip(test_df[SEQ_ID_COL].astype(str), test_df[SEQ_COL].astype(str)))\n    meta_by_id = {str(r[SEQ_ID_COL]): (r[STOI_COL] if STOI_COL else None, r[ALL_SEQ_COL] if ALL_SEQ_COL else None) for _, r in test_df.iterrows()}\n    protenix_preds = {}\n\n    for i in tqdm(range(len(dataset)), desc='Protenix'):\n        data, atom_array, error_message = dataset[i]\n        sample_name = data.get('sample_name', f'sample_{i}')\n        target_id = sample_name.split('_chunk')[0] if '_chunk' in sample_name else sample_name\n        if target_id not in seq_by_id:\n            continue\n\n        if error_message:\n            print(f\"  {target_id}: data error — {error_message}\")\n            continue\n\n        seq = seq_by_id[target_id]\n        stoi, all_seqs = meta_by_id[target_id]\n        chain_lens = get_chain_lengths(stoi, all_seqs, len(seq))\n\n        try:\n            new_cfg = update_inference_configs(configs, data['N_token'].item())\n            new_cfg.sample_diffusion.N_sample = N_SAMPLE\n            runner.update_model_configs(new_cfg)\n\n            prediction = runner.predict(data)\n            raw_coords = prediction['coordinate']\n            coords = extract_c1_coords(raw_coords, data, len(seq))\n\n            refined = []\n            for sample_idx in range(coords.shape[0]):\n                refined.append(refine_chain_geometry(coords[sample_idx], chain_lens, passes=3))\n            protenix_preds[target_id] = np.stack(refined, axis=0)\n            print(f\"  {target_id}: {protenix_preds[target_id].shape[0]} Protenix predictions ✓\")\n\n        except Exception as exc:\n            print(f\"  {target_id}: Protenix FAILED — {exc}\")\n            protenix_preds[target_id] = None\n\n        finally:\n            try:\n                del data, atom_array, error_message\n            except Exception:\n                pass\n            gc.collect()\n            if torch.cuda.is_available():\n                torch.cuda.empty_cache()\n\n    return protenix_preds\n","metadata":{"execution":{"iopub.execute_input":"2026-03-26T00:32:28.540875Z","iopub.status.busy":"2026-03-26T00:32:28.540606Z","iopub.status.idle":"2026-03-26T00:32:28.562217Z","shell.execute_reply":"2026-03-26T00:32:28.561544Z"},"language":"python","papermill":{"duration":0.042303,"end_time":"2026-03-26T00:32:28.563542","exception":false,"start_time":"2026-03-26T00:32:28.521239","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"3ae02d43","cell_type":"code","source":"predictions_store = {}\nprediction_labels_store = {}\n\nmodel_gpu1 = globals().get('model_gpu1', None)\n\nprotenix_predictions = protenix_predictions_for_test(test_seq)\nif protenix_predictions:\n    predictions_store.update(protenix_predictions)\n    print(f\"Loaded Protenix predictions for {len(protenix_predictions)} sequences.\")\n\nprint(f\"Generating predictions for {len(test_seq)} test sequences ...\")\nfor _, row in test_seq.iterrows():\n    sid = row[SEQ_ID_COL]\n    if sid in predictions_store:\n        continue\n    seq = row[SEQ_COL]\n    stoi = row[STOI_COL] if STOI_COL else None\n    all_seqs = row[ALL_SEQ_COL] if ALL_SEQ_COL else None\n    preds, labels = predict_5_diverse(\n        model, sid, seq, stoichiometry=stoi, all_sequences=all_seqs,\n        model_gpu1=model_gpu1)\n    predictions_store[sid] = preds\n    prediction_labels_store[sid] = labels\n\nprint(f\"Done. {len(predictions_store)} sequences x 5 predictions each.\")\n\nfirst_id = list(predictions_store.keys())[0]\nfirst_prd = predictions_store[first_id]\nfirst_lbl = prediction_labels_store[first_id]\nprint(f\"\\nSample - '{first_id}' (L={len(first_prd[0])}):\\n\")\nfor i, p in enumerate(first_prd):\n    step = np.linalg.norm(np.diff(p, axis=0), axis=1).mean() if len(p) > 1 else float('nan')\n    tag = first_lbl[i] if i < len(first_lbl) else f'pred{i+1}'\n    print(f\"  pred {i+1} ({tag:12s}): \"\n          f\"x=[{p[:,0].min():.1f},{p[:,0].max():.1f}]  \"\n          f\"y=[{p[:,1].min():.1f},{p[:,1].max():.1f}]  \"\n          f\"z=[{p[:,2].min():.1f},{p[:,2].max():.1f}]  \"\n          f\"mean_step={step:.2f}A\")","metadata":{"execution":{"iopub.execute_input":"2026-03-26T00:32:28.596545Z","iopub.status.busy":"2026-03-26T00:32:28.596226Z","iopub.status.idle":"2026-03-26T00:34:25.266723Z","shell.execute_reply":"2026-03-26T00:34:25.265881Z"},"papermill":{"duration":116.702494,"end_time":"2026-03-26T00:34:25.283417","exception":false,"start_time":"2026-03-26T00:32:28.580923","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"0c5c654e","cell_type":"markdown","source":"## 7. Create Submission\n\nOne row per residue. 5 sets of x,y,z coordinates per row.\n\n| Column | Value |\n|---|---|\n| `ID` | `{target_id}_{resid}` |\n| `resname` | Single-letter nucleotide |\n| `resid` | 1-based residue index |\n| `x_1 … z_5` | Coordinates for each prediction (Å, 3dp) |","metadata":{"papermill":{"duration":0.014678,"end_time":"2026-03-26T00:34:25.312607","exception":false,"start_time":"2026-03-26T00:34:25.297929","status":"completed"},"tags":[]}},{"id":"5bb0b999","cell_type":"code","source":"rows = []\n\nfor _, row in test_seq.iterrows():\n    sid = row[SEQ_ID_COL]\n    sequence = row[SEQ_COL]\n    preds = predictions_store[sid]\n\n    for res_idx, resname in enumerate(sequence, start=1):\n        entry = {\n            'ID': f\"{sid}_{res_idx}\",\n            'resname': resname,\n            'resid': res_idx,\n        }\n        for pred_num, pred_coords in enumerate(preds, start=1):\n            x, y, z = pred_coords[res_idx - 1]\n            entry[f'x_{pred_num}'] = round(float(x), 3)\n            entry[f'y_{pred_num}'] = round(float(y), 3)\n            entry[f'z_{pred_num}'] = round(float(z), 3)\n        rows.append(entry)\n\nsubmission = pd.DataFrame(rows)\nexpected_cols = ['ID', 'resname', 'resid', 'x_1', 'y_1', 'z_1', 'x_2', 'y_2', 'z_2', 'x_3', 'y_3', 'z_3', 'x_4', 'y_4', 'z_4', 'x_5', 'y_5', 'z_5']\nsubmission = submission.reindex(columns=expected_cols)\n\nnum_cols = [c for c in expected_cols if c not in ('ID', 'resname', 'resid')]\nsubmission[num_cols] = submission[num_cols].apply(pd.to_numeric, errors='coerce')\nmissing = [c for c in expected_cols if c not in submission.columns]\nif missing:\n    print(f\"⚠️  Missing columns: {missing}\")\n\ninvalid_mask = ~np.isfinite(submission[num_cols].to_numpy()).all(axis=1)\nif invalid_mask.any():\n    bad_ids = submission.loc[invalid_mask, 'ID'].head(10).tolist()\n    print(f\"⚠️  Non-finite values remain for IDs like: {bad_ids}\")\n\nprint(\"✓  Column order matches sample submission.\")\nsubmission.to_csv(SUBMISSION_FILE, index=False)\nprint(f\"✓  Submission saved → {SUBMISSION_FILE}\")\nprint(f\"   Shape: {submission.shape}\")\ndisplay(submission.head(10))","metadata":{"execution":{"iopub.execute_input":"2026-03-26T00:34:25.345356Z","iopub.status.busy":"2026-03-26T00:34:25.344949Z","iopub.status.idle":"2026-03-26T00:34:25.722206Z","shell.execute_reply":"2026-03-26T00:34:25.721599Z"},"papermill":{"duration":0.394777,"end_time":"2026-03-26T00:34:25.723761","exception":false,"start_time":"2026-03-26T00:34:25.328984","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"a7f39251","cell_type":"markdown","source":"## References\n\n### Primary Architecture\n\n1. **RhoFold+** — Shen, T. et al. (2024). *Accurate RNA 3D structure prediction using a language model-based deep learning approach.* Nature Methods, 21, 2287–2298.\n   https://doi.org/10.1038/s41592-024-02487-0\n   → RNA-FM embeddings, IPA structure module, MSA subsampling diversity, recycling.\n\n2. **AlphaFold2** — Jumper, J. et al. (2021). *Highly accurate protein structure prediction with AlphaFold.* Nature, 596, 583–589.\n   https://doi.org/10.1038/s41586-021-03819-2\n   → Invariant Point Attention (IPA), FAPE loss, Evoformer pair track, recycling.\n\n3. **trRosettaRNA** — Wang, W. et al. (2023). *trRosettaRNA: automated prediction of RNA 3D structure with transformer network.* Nature Communications, 14, 7266.\n   https://doi.org/10.1038/s41467-023-42528-4\n   → Relative Position Bias (RPB), MSA coevolution features.\n\n4. **DeepFoldRNA** — Pearce, R. et al. (2022). *De novo RNA tertiary structure prediction at atomic resolution using geometric potentials from deep learning.* bioRxiv.\n   https://doi.org/10.1101/2022.05.15.491755\n   → MI coevolution features, distance distogram loss.\n\n### Pre-trained RNA Language Models\n\n5. **RNA-FM** — Chen, J. et al. (2022). *Interpretable RNA Foundation Model from Unannotated Data for Highly Accurate RNA Structure and Function Predictions.* arXiv:2204.00300.\n   → Pre-trained nucleotide-level embeddings (640-dim) used as frozen backbone.\n\n### Reviews\n\n6. **RNA 3D Structure Prediction — Molecules Review** — Wang, X. et al. (2023). *RNA 3D structure prediction: progress and perspective.* Molecules, 28(14), 5532.\n   https://doi.org/10.3390/molecules28145532\n\n7. **RNA 3D Structure Modeling — Frontiers Review** — Li, B. et al. (2020). *Advances in RNA 3D structure modeling using experimental data.* Frontiers in Genetics, 11, 574485.\n   https://doi.org/10.3389/fgene.2020.574485\n\n### Techniques Implemented\n\n| Technique | Source |\n|---|---|\n| Frozen RNA-FM embeddings (640-dim) | RNA-FM, RhoFold+ |\n| Pair track (outer product mean + triangle attention/multiplication) | AlphaFold2 Evoformer |\n| IPA structure module | AlphaFold2, RhoFold+ |\n| FAPE loss | AlphaFold2 |\n| Distance distogram CE loss | DeepFoldRNA, AlphaFold2 |\n| RPB in single-track attention | trRosettaRNA |\n| MI coevolution features | DeepFoldRNA, trRosettaRNA |\n| Recycling (4 iterations) | AlphaFold2, RhoFold+ |\n| MSA subsampling diversity | RhoFold+ |\n| EMA weight averaging | Standard practice |","metadata":{"papermill":{"duration":0.014576,"end_time":"2026-03-26T00:34:25.754855","exception":false,"start_time":"2026-03-26T00:34:25.740279","status":"completed"},"tags":[]}}]}