{"cells":[{"cell_type":"markdown","metadata":{},"source":"# RNA 3D Folding Part 2 — Submission Notebook\n**Hybrid family-aware pipeline:**\nRNA sequence → secondary structure → family classification → template search\n→ TBM/de-novo routing → Protenix inference → motif correction (GNRA, K-turn)\n→ 5-candidate ensemble → submission.csv\n\nStrategy based on:\n- **Approach A** (sigmaborov, LB 0.438): Protenix + RibonanzaNet2 baseline\n- **Approach B** (gourabr0y555): Protenix + TBM templates\n- **Approach C** (artemevstafyev): High-score without hash tricks (pure DL)\n","outputs":[]},{"cell_type":"code","metadata":{},"source":"# ────────────────────────────────────────────────────────────\n# Environment Setup\n# ────────────────────────────────────────────────────────────\nimport os\nimport sys\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\n# ── Environment detection ─────────────────────────────────────────────\nif os.path.exists(\"/kaggle/input\"):\n    # Running on Kaggle\n    KAGGLE_ENV   = True\n    DATA_DIR     = \"/kaggle/input/stanford-rna-3d-folding-2\"\n    OUTPUT_DIR   = \"/kaggle/working\"\n    # Pipeline src/ is injected below from the pipeline dataset\n    PIPELINE_DIR = \"/kaggle/input/rna-pipeline-src\"\nelse:\n    # Running locally\n    KAGGLE_ENV   = False\n    DATA_DIR     = \"/home/ilan/kaggle/data\"\n    OUTPUT_DIR   = \".\"\n    PIPELINE_DIR = \".\"\n\n# ── Add src/ to path ──────────────────────────────────────────────────\n# On Kaggle: add PIPELINE_DIR (dataset containing repo src/)\n# Locally:   add current dir (repo root)\nsys.path.insert(0, PIPELINE_DIR)\n\nprint(f\"Environment : {'KAGGLE' if KAGGLE_ENV else 'LOCAL'}\")\nprint(f\"DATA_DIR    : {DATA_DIR}\")\nprint(f\"OUTPUT_DIR  : {OUTPUT_DIR}\")\nprint(f\"Python      : {sys.version.split()[0]}\")\n","outputs":[],"execution_count":null},{"cell_type":"code","metadata":{},"source":"# ────────────────────────────────────────────────────────────\n# Dependency Check\n# ────────────────────────────────────────────────────────────\n# Install bioinformatics tools needed by the pipeline\n# (These are pre-installed in many Kaggle environments; safe to run regardless)\nimport subprocess\n\ndef install_if_missing(cmd_check: str, install_cmd: list):\n    result = subprocess.run([\"which\", cmd_check], capture_output=True)\n    if result.returncode != 0:\n        print(f\"Installing {cmd_check}...\")\n        subprocess.run(install_cmd, check=False, capture_output=True)\n    else:\n        print(f\"  {cmd_check}: already available\")\n\n# ViennaRNA (secondary structure)\ntry:\n    import RNA\n    print(\"  ViennaRNA Python: available\")\nexcept ImportError:\n    print(\"  ViennaRNA Python: not available, using subprocess RNAfold\")\n\n# Check command-line tools\nfor tool, apt_pkg in [\n    (\"RNAfold\",  \"vienna-rna\"),\n    (\"mmseqs\",   \"\"),   # install from bioconda via setup.sh\n    (\"cmscan\",   \"infernal\"),\n    (\"USalign\",  \"\"),   # compiled in setup.sh\n]:\n    result = subprocess.run([\"which\", tool], capture_output=True)\n    status = \"✓ available\" if result.returncode == 0 else \"✗ not found (fallback active)\"\n    print(f\"  {tool:12s}: {status}\")\n\nprint(\"\\nNote: missing tools trigger graceful fallbacks in the pipeline.\")\n","outputs":[],"execution_count":null},{"cell_type":"code","metadata":{},"source":"# ────────────────────────────────────────────────────────────\n# Load Competition Data\n# ────────────────────────────────────────────────────────────\nimport pandas as pd\nimport numpy as np\nfrom pathlib import Path\n\n# Load test sequences\ntest_sequences = pd.read_csv(f\"{DATA_DIR}/test_sequences.csv\")\ntest_sequences[\"sequence\"] = test_sequences[\"sequence\"].str.upper().str.replace(\"T\", \"U\")\n\n# Parse stoichiometry metadata\ntest_sequences[\"n_copies\"]   = test_sequences[\"stoichiometry\"].str.extract(r\":(\\d+)$\").astype(float).fillna(1).astype(int)\ntest_sequences[\"has_ligands\"] = test_sequences[\"ligand_ids\"].notna() & (test_sequences[\"ligand_ids\"].astype(str).str.len() > 1)\ntest_sequences[\"is_complex\"] = test_sequences[\"stoichiometry\"].str.contains(\";\") | (test_sequences[\"n_copies\"] > 1)\ntest_sequences[\"seq_len\"]    = test_sequences[\"sequence\"].str.len()\n\nprint(f\"Test sequences loaded : {len(test_sequences)}\")\nprint(f\"Length range          : {test_sequences['seq_len'].min()} – {test_sequences['seq_len'].max()} nt\")\nprint(f\"Has ligands           : {test_sequences['has_ligands'].sum()} / {len(test_sequences)}\")\nprint(f\"Multi-chain           : {test_sequences['is_complex'].sum()} / {len(test_sequences)}\")\nprint()\nprint(test_sequences[[\"target_id\",\"seq_len\",\"stoichiometry\",\"has_ligands\"]].to_string())\n","outputs":[],"execution_count":null},{"cell_type":"code","metadata":{},"source":"# ────────────────────────────────────────────────────────────\n# Utils — sequence helpers\n# ────────────────────────────────────────────────────────────\n\"\"\"utils/sequence_utils.py — RNA sequence utilities.\"\"\"\n\nimport re\n\nVALID_RNA_CHARS = set(\"ACGUacguNn\")\nIUPAC_MAP = {\n    \"A\": \"A\", \"C\": \"C\", \"G\": \"G\", \"U\": \"U\", \"T\": \"U\",\n    \"R\": \"AG\", \"Y\": \"CU\", \"S\": \"GC\", \"W\": \"AU\",\n    \"K\": \"GU\", \"M\": \"AC\", \"B\": \"CGU\", \"D\": \"AGU\",\n    \"H\": \"ACU\", \"V\": \"ACG\", \"N\": \"ACGU\",\n}\n\n\ndef validate_rna_sequence(sequence: str) -> bool:\n    \"\"\"Return True if sequence contains only valid RNA/IUPAC characters.\"\"\"\n    return bool(sequence) and all(c.upper() in IUPAC_MAP for c in sequence)\n\n\ndef normalize_sequence(sequence: str) -> str:\n    \"\"\"Uppercase and replace T→U.\"\"\"\n    return sequence.upper().replace(\"T\", \"U\")\n\n\ndef gc_content(sequence: str) -> float:\n    \"\"\"Return GC fraction of the sequence.\"\"\"\n    seq = normalize_sequence(sequence)\n    gc = sum(1 for c in seq if c in (\"G\", \"C\"))\n    return gc / len(seq) if seq else 0.0\n\n\ndef split_into_chunks(sequence: str, chunk_size: int, overlap: int = 50) -> list[tuple[int, int]]:\n    \"\"\"\n    Split a long sequence into overlapping chunks.\n    Returns list of (start, end) tuples (0-indexed, end exclusive).\n    \"\"\"\n    n = len(sequence)\n    if n <= chunk_size:\n        return [(0, n)]\n    chunks = []\n    pos = 0\n    while pos < n:\n        end = min(pos + chunk_size, n)\n        chunks.append((pos, end))\n        if end == n:\n            break\n        pos += chunk_size - overlap\n    return chunks\n\n\ndef format_fasta(target_id: str, sequence: str) -> str:\n    \"\"\"Format a sequence as FASTA.\"\"\"\n    return f\">{target_id}\\n{sequence}\\n\"\n","outputs":[],"execution_count":null},{"cell_type":"code","metadata":{},"source":"# ────────────────────────────────────────────────────────────\n# Stage 1 — Secondary structure prediction\n# ────────────────────────────────────────────────────────────\n\"\"\"\nsecondary_structure.py — Secondary structure prediction wrapper.\n\nSupports: ViennaRNA (RNAfold), EternaFold, CONTRAfold\nOutput: SecondaryStructure dataclass with dot-bracket, base pairs, stems, loops\n\"\"\"\n\nimport subprocess\nimport re\nfrom dataclasses import dataclass, field\nfrom typing import Optional\nimport logging\n\nlogger = logging.getLogger(__name__)\n\n\n@dataclass\nclass StemLoop:\n    \"\"\"A stem-loop (hairpin) structural element.\"\"\"\n    stem_start: int\n    stem_end: int\n    loop_start: int\n    loop_end: int\n    loop_sequence: str\n\n    @property\n    def loop_length(self) -> int:\n        return self.loop_end - self.loop_start + 1\n\n\n@dataclass\nclass SecondaryStructure:\n    \"\"\"Container for secondary structure prediction results.\"\"\"\n    sequence: str\n    dot_bracket: str\n    mfe: float                              # minimum free energy (kcal/mol)\n    base_pairs: list[tuple[int, int]]       # (i, j) 0-indexed\n    stems: list[tuple[int, int, int, int]]  # (i_start, i_end, j_start, j_end)\n    hairpins: list[StemLoop]\n    engine: str = \"viennarna\"\n\n    # Detected motif positions (filled by MotifCorrector)\n    gnra_positions: list[int] = field(default_factory=list)\n    kturn_positions: list[int] = field(default_factory=list)\n\n    @property\n    def n_pairs(self) -> int:\n        return len(self.base_pairs)\n\n    @property\n    def pair_fraction(self) -> float:\n        return (2 * self.n_pairs) / len(self.sequence) if self.sequence else 0.0\n\n    def has_hairpin_of_length(self, n: int) -> bool:\n        return any(h.loop_length == n for h in self.hairpins)\n\n\nclass SecondaryStructurePredictor:\n    \"\"\"\n    Wraps ViennaRNA (RNAfold) as the default engine.\n    Falls back gracefully if not installed (returns minimal structure).\n    \"\"\"\n\n    def __init__(self, cfg: dict):\n        self.engine = cfg.get(\"engine\", \"viennarna\")\n        self.temperature = cfg.get(\"temperature\", 37.0)\n        self.use_pseudoknot = cfg.get(\"use_pseudoknot\", False)\n        self._viennarna_available = self._check_viennarna()\n\n    def _check_viennarna(self) -> bool:\n        try:\n            result = subprocess.run(\n                [\"RNAfold\", \"--version\"],\n                capture_output=True, text=True, timeout=5\n            )\n            return result.returncode == 0\n        except FileNotFoundError:\n            logger.warning(\"RNAfold not found. Using simple bracket fallback.\")\n            return False\n\n    def predict(self, sequence: str) -> SecondaryStructure:\n        \"\"\"Predict secondary structure for a given RNA sequence.\"\"\"\n        if self.engine == \"viennarna\" and self._viennarna_available:\n            return self._predict_viennarna(sequence)\n        else:\n            logger.warning(f\"Engine '{self.engine}' unavailable. Using fallback.\")\n            return self._predict_fallback(sequence)\n\n    def _predict_viennarna(self, sequence: str) -> SecondaryStructure:\n        \"\"\"Call RNAfold subprocess and parse output.\"\"\"\n        cmd = [\"RNAfold\", \"--noPS\", f\"--temp={self.temperature}\"]\n        if self.use_pseudoknot:\n            cmd.append(\"--gquad\")  # G-quadruplex support\n\n        result = subprocess.run(\n            cmd,\n            input=sequence,\n            capture_output=True, text=True, timeout=120\n        )\n        if result.returncode != 0:\n            logger.error(f\"RNAfold failed: {result.stderr}\")\n            return self._predict_fallback(sequence)\n\n        lines = result.stdout.strip().split(\"\\n\")\n        # ViennaRNA output: line0 = sequence, line1 = \"dotbracket (MFE)\"\n        db_line = lines[-1]\n        # Parse: \"(((...))) (-12.34)\"\n        match = re.match(r'^([.()[\\]{}]+)\\s+\\((-?[\\d.]+)\\)', db_line)\n        if not match:\n            logger.warning(f\"Could not parse RNAfold output: {db_line}\")\n            return self._predict_fallback(sequence)\n\n        dot_bracket = match.group(1)\n        mfe = float(match.group(2))\n\n        base_pairs = self._parse_base_pairs(dot_bracket)\n        stems = self._extract_stems(base_pairs)\n        hairpins = self._extract_hairpins(sequence, dot_bracket, base_pairs)\n\n        return SecondaryStructure(\n            sequence=sequence,\n            dot_bracket=dot_bracket,\n            mfe=mfe,\n            base_pairs=base_pairs,\n            stems=stems,\n            hairpins=hairpins,\n            engine=\"viennarna\",\n        )\n\n    def _predict_fallback(self, sequence: str) -> SecondaryStructure:\n        \"\"\"Minimal fallback: return fully unstructured.\"\"\"\n        n = len(sequence)\n        return SecondaryStructure(\n            sequence=sequence,\n            dot_bracket=\".\" * n,\n            mfe=0.0,\n            base_pairs=[],\n            stems=[],\n            hairpins=[],\n            engine=\"fallback\",\n        )\n\n    def _parse_base_pairs(self, dot_bracket: str) -> list[tuple[int, int]]:\n        \"\"\"Parse dot-bracket notation into list of (i,j) base pairs (0-indexed).\"\"\"\n        pairs = []\n        stack = []\n        for i, c in enumerate(dot_bracket):\n            if c == \"(\":\n                stack.append(i)\n            elif c == \")\":\n                if stack:\n                    j = stack.pop()\n                    pairs.append((j, i))\n        return sorted(pairs)\n\n    def _extract_stems(\n        self, base_pairs: list[tuple[int, int]]\n    ) -> list[tuple[int, int, int, int]]:\n        \"\"\"\n        Group consecutive base pairs into stems.\n        A stem is a run of pairs (i, j), (i+1, j-1), (i+2, j-2), ...\n        Returns list of (i_start, i_end, j_start, j_end).\n        \"\"\"\n        if not base_pairs:\n            return []\n        stems = []\n        current_stem = [base_pairs[0]]\n        for prev, curr in zip(base_pairs, base_pairs[1:]):\n            if curr[0] == prev[0] + 1 and curr[1] == prev[1] - 1:\n                current_stem.append(curr)\n            else:\n                if len(current_stem) >= 2:\n                    stems.append((\n                        current_stem[0][0], current_stem[-1][0],\n                        current_stem[-1][1], current_stem[0][1],\n                    ))\n                current_stem = [curr]\n        if len(current_stem) >= 2:\n            stems.append((\n                current_stem[0][0], current_stem[-1][0],\n                current_stem[-1][1], current_stem[0][1],\n            ))\n        return stems\n\n    def _extract_hairpins(\n        self,\n        sequence: str,\n        dot_bracket: str,\n        base_pairs: list[tuple[int, int]],\n    ) -> list[StemLoop]:\n        \"\"\"Extract hairpin loops (closing pair + unpaired loop region).\"\"\"\n        hairpins = []\n        pair_set = set(base_pairs)\n        for i, j in base_pairs:\n            # Check if all residues between i and j are unpaired\n            inner = dot_bracket[i + 1:j]\n            if all(c == \".\" for c in inner):\n                loop_seq = sequence[i + 1:j]\n                hairpins.append(StemLoop(\n                    stem_start=i,\n                    stem_end=j,\n                    loop_start=i + 1,\n                    loop_end=j - 1,\n                    loop_sequence=loop_seq,\n                ))\n        return hairpins\n","outputs":[],"execution_count":null},{"cell_type":"code","metadata":{},"source":"# ────────────────────────────────────────────────────────────\n# Stage 2 — Family classification (Rfam heuristic)\n# ────────────────────────────────────────────────────────────\n\"\"\"\nfamily_classifier.py — Classify RNA sequences into structural families.\n\nUses Infernal (cmscan) against the Rfam covariance model database.\nFalls back to sequence-pattern heuristics if Infernal is unavailable.\n\nFamilies relevant to this competition:\n  - riboswitch     (aptamer domain + expression platform)\n  - tRNA           (cloverleaf, universal template)\n  - ribosomal      (rRNA fragments)\n  - viral          (IRES, frameshifting elements, etc.)\n  - aptamer        (synthetic / in-vitro selected)\n  - ribozyme       (catalytic RNA)\n  - unknown        → force de novo branch\n\"\"\"\n\nimport subprocess\nimport re\nimport logging\nfrom dataclasses import dataclass\nfrom pathlib import Path\nfrom typing import Optional\n\nfrom src.secondary_structure import SecondaryStructure\n\nlogger = logging.getLogger(__name__)\n\n# Rfam family ID → human-readable category mapping\nRFAM_CATEGORY = {\n    # Riboswitches\n    \"RF00050\": \"riboswitch\",  # FMN\n    \"RF00059\": \"riboswitch\",  # TPP\n    \"RF00162\": \"riboswitch\",  # SAM-I\n    \"RF00174\": \"riboswitch\",  # Cobalamin\n    \"RF00234\": \"riboswitch\",  # glmS\n    \"RF01057\": \"riboswitch\",  # SAH\n    \"RF01786\": \"riboswitch\",  # ZMP/ZTP\n    # tRNA / tRNA-like\n    \"RF00005\": \"tRNA\",\n    \"RF00023\": \"tRNA\",        # tmRNA\n    # Ribosomal RNA\n    \"RF00177\": \"ribosomal\",   # SSU rRNA\n    \"RF02540\": \"ribosomal\",   # LSU rRNA\n    \"RF01960\": \"ribosomal\",   # SSU rRNA archaea\n    # Ribozymes\n    \"RF00008\": \"ribozyme\",    # Hammerhead type III\n    \"RF00163\": \"ribozyme\",    # Hammerhead type I\n    \"RF00622\": \"ribozyme\",    # HDV\n    # Viral\n    \"RF00164\": \"viral\",       # Corona 3' UTR\n    \"RF00165\": \"viral\",       # Corona 5' UTR\n    # Large ncRNA\n    \"RF02348\": \"large_ncrna\", # OLE RNA\n    \"RF02357\": \"large_ncrna\", # GOLLD RNA\n    \"RF02544\": \"large_ncrna\", # ROOL RNA\n}\n\n\n@dataclass\nclass FamilyResult:\n    \"\"\"Result of family classification.\"\"\"\n    name: str               # e.g. \"riboswitch\", \"tRNA\", \"unknown\"\n    rfam_id: Optional[str]  # e.g. \"RF00059\"\n    score: float            # bit score from cmscan, or heuristic confidence\n    evalue: float\n    is_known: bool          # True if a Rfam hit was found\n\n    def __str__(self):\n        return f\"{self.name}({self.rfam_id or 'heuristic'} e={self.evalue:.1e})\"\n\n\nclass FamilyClassifier:\n    \"\"\"\n    RNA family classifier.\n    Priority order:\n      1. Rfam cmscan (if Infernal installed + Rfam.cm available)\n      2. Sequence heuristics (motif patterns)\n      3. Unknown\n    \"\"\"\n\n    def __init__(self, cfg: dict):\n        self.rfam_db = cfg.get(\"rfam_db\", \"data/rfam/Rfam.cm\")\n        self.evalue_threshold = cfg.get(\"evalue_threshold\", 1e-5)\n        self.known_families = cfg.get(\"known_families\", [])\n        self._infernal_available = self._check_infernal()\n        self._rfam_available = Path(self.rfam_db).exists()\n        if not self._rfam_available:\n            logger.warning(f\"Rfam CM not found at {self.rfam_db}. Using heuristics.\")\n\n    def _check_infernal(self) -> bool:\n        try:\n            result = subprocess.run(\n                [\"cmscan\", \"-h\"], capture_output=True, text=True, timeout=5\n            )\n            return result.returncode == 0\n        except FileNotFoundError:\n            logger.warning(\"cmscan not found. Using sequence heuristics for family classification.\")\n            return False\n\n    def classify(self, sequence: str, sec_struct: SecondaryStructure) -> FamilyResult:\n        \"\"\"Classify an RNA sequence into a structural family.\"\"\"\n        if self._infernal_available and self._rfam_available:\n            result = self._classify_cmscan(sequence)\n            if result is not None:\n                return result\n\n        # Fallback to sequence/structure heuristics\n        return self._classify_heuristic(sequence, sec_struct)\n\n    def _classify_cmscan(self, sequence: str) -> Optional[FamilyResult]:\n        \"\"\"Run cmscan against Rfam and parse the best hit.\"\"\"\n        import tempfile, os\n        with tempfile.NamedTemporaryFile(mode=\"w\", suffix=\".fa\", delete=False) as f:\n            f.write(f\">query\\n{sequence}\\n\")\n            fasta_path = f.name\n        try:\n            result = subprocess.run(\n                [\n                    \"cmscan\", \"--tblout\", \"/dev/stdout\",\n                    \"-E\", str(self.evalue_threshold),\n                    \"--noali\", \"--cpu\", \"2\",\n                    self.rfam_db, fasta_path,\n                ],\n                capture_output=True, text=True, timeout=120\n            )\n        except subprocess.TimeoutExpired:\n            logger.warning(\"cmscan timed out\")\n            return None\n        finally:\n            os.unlink(fasta_path)\n\n        if result.returncode != 0:\n            logger.warning(f\"cmscan error: {result.stderr[:200]}\")\n            return None\n\n        # Parse tblout format\n        best_hit = None\n        for line in result.stdout.split(\"\\n\"):\n            if line.startswith(\"#\") or not line.strip():\n                continue\n            parts = line.split()\n            if len(parts) < 16:\n                continue\n            # tblout columns: target_name, accession, query, clan, mdl, mdl_from, mdl_to,\n            #                  seq_from, seq_to, strand, trunc, pass, gc, bias, score, E-value\n            try:\n                rfam_acc = parts[1]\n                score = float(parts[14])\n                evalue = float(parts[15])\n                if best_hit is None or evalue < best_hit[2]:\n                    best_hit = (rfam_acc, score, evalue)\n            except (ValueError, IndexError):\n                continue\n\n        if best_hit is None:\n            return None\n\n        rfam_id, score, evalue = best_hit\n        category = RFAM_CATEGORY.get(rfam_id, \"other_ncrna\")\n        return FamilyResult(\n            name=category,\n            rfam_id=rfam_id,\n            score=score,\n            evalue=evalue,\n            is_known=True,\n        )\n\n    def _classify_heuristic(\n        self, sequence: str, sec_struct: SecondaryStructure\n    ) -> FamilyResult:\n        \"\"\"\n        Fast heuristic classification based on:\n        - Sequence length\n        - Nucleotide composition\n        - Hairpin count and sizes\n        - Known motif patterns\n        \"\"\"\n        n = len(sequence)\n\n        # tRNA heuristic: ~73-93 nt, cloverleaf structure (4 stems)\n        if 70 <= n <= 100 and len(sec_struct.stems) >= 3:\n            return FamilyResult(\"tRNA\", None, 0.6, 0.1, False)\n\n        # Short ribozyme heuristic: <60 nt, highly structured\n        if n < 70 and sec_struct.pair_fraction > 0.5:\n            return FamilyResult(\"ribozyme\", None, 0.4, 0.5, False)\n\n        # Riboswitch-like: 80-300 nt, multiple stems\n        if 80 <= n <= 350 and len(sec_struct.stems) >= 2:\n            return FamilyResult(\"riboswitch\", None, 0.3, 1.0, False)\n\n        # Large ncRNA (Part 2 specific): >300 nt\n        if n > 300:\n            return FamilyResult(\"large_ncrna\", None, 0.2, 5.0, False)\n\n        return FamilyResult(\"unknown\", None, 0.0, 999.0, False)\n","outputs":[],"execution_count":null},{"cell_type":"code","metadata":{},"source":"# ────────────────────────────────────────────────────────────\n# Stage 3 — Template search (MMseqs2 / PDB cache)\n# ────────────────────────────────────────────────────────────\n\"\"\"\ntemplate_search.py — Template-Based Modeling (TBM) search module.\n\nSearches a local PDB RNA C1' coordinate database using MMseqs2.\nReturns ranked structural templates with expected TM-score estimates.\n\nBased on the approach from:\n  - Notebook B (gourabr0y555): Protenix + TBM\n  - jaejohn (Part 1 1st place): TBM-only approach\n  - NVIDIA RNAPro: MMseqs2 3D RNA Template Identification\n\"\"\"\n\nimport subprocess\nimport pickle\nimport logging\nimport tempfile\nimport os\nfrom dataclasses import dataclass, field\nfrom pathlib import Path\nfrom typing import Optional\n\nimport numpy as np\n\nfrom src.family_classifier import FamilyResult\n\nlogger = logging.getLogger(__name__)\n\n\n@dataclass\nclass Template:\n    \"\"\"A single structural template retrieved from PDB.\"\"\"\n    pdb_id: str\n    chain_id: str\n    sequence: str\n    seq_identity: float       # fraction (0–1)\n    coverage: float           # query coverage (0–1)\n    expected_tm: float        # estimated TM-score for this template\n    c1_coords: Optional[np.ndarray] = None  # shape (L, 3)\n    alignment: dict = field(default_factory=dict)\n\n    @property\n    def label(self) -> str:\n        return f\"{self.pdb_id}_{self.chain_id}\"\n\n    def to_dict(self) -> dict:\n        return {\n            \"pdb_id\": self.pdb_id,\n            \"chain_id\": self.chain_id,\n            \"sequence\": self.sequence,\n            \"seq_identity\": self.seq_identity,\n            \"coverage\": self.coverage,\n            \"expected_tm\": self.expected_tm,\n            \"c1_coords\": self.c1_coords,\n        }\n\n\nclass TemplateSearcher:\n    \"\"\"\n    Two-stage template search:\n      1. MMseqs2 for fast sequence similarity search against PDB RNA sequences\n      2. Load C1' coordinates for top hits from local cache\n    \"\"\"\n\n    def __init__(self, cfg: dict):\n        self.enabled = cfg.get(\"enabled\", True)\n        self.mmseqs2_db = cfg.get(\"mmseqs2_db\", \"data/pdb_cache/pdb_rna_mmseqs2\")\n        self.pdb_c1_cache = cfg.get(\"pdb_c1_cache\", \"data/pdb_cache/pdb_c1_coords.pkl\")\n        self.max_templates = cfg.get(\"max_templates\", 10)\n        self.min_seq_identity = cfg.get(\"min_seq_identity\", 0.25)\n        self.min_coverage = cfg.get(\"min_coverage\", 0.5)\n\n        self._mmseqs2_available = self._check_mmseqs2()\n        self._c1_cache = self._load_c1_cache()\n\n    def _check_mmseqs2(self) -> bool:\n        try:\n            r = subprocess.run([\"mmseqs\", \"version\"], capture_output=True, timeout=5)\n            return r.returncode == 0\n        except FileNotFoundError:\n            logger.warning(\"MMseqs2 not found. Template search disabled.\")\n            return False\n\n    def _load_c1_cache(self) -> dict:\n        \"\"\"Load prebuilt C1' coordinate cache from disk.\"\"\"\n        cache_path = Path(self.pdb_c1_cache)\n        if cache_path.exists():\n            logger.info(f\"Loading C1' cache from {cache_path}\")\n            with open(cache_path, \"rb\") as f:\n                return pickle.load(f)\n        else:\n            logger.warning(f\"C1' cache not found at {cache_path}. Run build_pdb_cache.sh first.\")\n            return {}\n\n    def search(self, sequence: str, family: FamilyResult) -> list[Template]:\n        \"\"\"\n        Search for structural templates for the given sequence.\n        Returns list of Template objects sorted by expected TM-score descending.\n        \"\"\"\n        if not self.enabled:\n            return []\n        if not self._mmseqs2_available:\n            return []\n        if not Path(self.mmseqs2_db).exists():\n            logger.warning(f\"MMseqs2 DB not found at {self.mmseqs2_db}\")\n            return []\n\n        raw_hits = self._run_mmseqs2(sequence)\n        templates = []\n\n        for hit in raw_hits[:self.max_templates * 3]:  # search wider, filter after\n            if hit[\"seq_identity\"] < self.min_seq_identity:\n                continue\n            if hit[\"coverage\"] < self.min_coverage:\n                continue\n\n            # Estimate TM-score from sequence identity (empirical formula)\n            expected_tm = self._estimate_tm_from_seqid(\n                hit[\"seq_identity\"], hit[\"coverage\"], len(sequence)\n            )\n\n            # Load C1' coordinates from cache\n            key = f\"{hit['pdb_id']}_{hit['chain_id']}\"\n            c1_coords = self._c1_cache.get(key)\n\n            templates.append(Template(\n                pdb_id=hit[\"pdb_id\"],\n                chain_id=hit[\"chain_id\"],\n                sequence=hit[\"target_seq\"],\n                seq_identity=hit[\"seq_identity\"],\n                coverage=hit[\"coverage\"],\n                expected_tm=expected_tm,\n                c1_coords=c1_coords,\n                alignment=hit.get(\"alignment\", {}),\n            ))\n\n        # Sort by expected TM-score\n        templates.sort(key=lambda t: t.expected_tm, reverse=True)\n        return templates[:self.max_templates]\n\n    def _run_mmseqs2(self, sequence: str) -> list[dict]:\n        \"\"\"Run MMseqs2 easy-search and parse hits.\"\"\"\n        with tempfile.TemporaryDirectory() as tmpdir:\n            query_fa = os.path.join(tmpdir, \"query.fa\")\n            result_tsv = os.path.join(tmpdir, \"result.tsv\")\n            tmp_mmseqs = os.path.join(tmpdir, \"tmp\")\n\n            with open(query_fa, \"w\") as f:\n                f.write(f\">query\\n{sequence}\\n\")\n\n            cmd = [\n                \"mmseqs\", \"easy-search\",\n                query_fa, self.mmseqs2_db, result_tsv, tmp_mmseqs,\n                \"--format-output\",\n                \"query,target,pident,alnlen,qstart,qend,tstart,tend,evalue,bits,qaln,taln\",\n                \"--min-seq-id\", str(self.min_seq_identity),\n                \"-c\", str(self.min_coverage),\n                \"--cov-mode\", \"0\",\n                \"-e\", \"0.001\",\n                \"--threads\", \"4\",\n                \"-v\", \"0\",\n            ]\n\n            try:\n                result = subprocess.run(\n                    cmd, capture_output=True, text=True, timeout=120\n                )\n            except subprocess.TimeoutExpired:\n                logger.warning(\"MMseqs2 search timed out\")\n                return []\n\n            if result.returncode != 0:\n                logger.error(f\"MMseqs2 failed: {result.stderr[:300]}\")\n                return []\n\n            return self._parse_mmseqs2_output(result_tsv, len(sequence))\n\n    def _parse_mmseqs2_output(self, tsv_path: str, query_len: int) -> list[dict]:\n        \"\"\"Parse MMseqs2 tabular output.\"\"\"\n        hits = []\n        if not os.path.exists(tsv_path):\n            return hits\n        with open(tsv_path) as f:\n            for line in f:\n                parts = line.strip().split(\"\\t\")\n                if len(parts) < 10:\n                    continue\n                try:\n                    target = parts[1]  # e.g. \"4XWF_A\"\n                    pident = float(parts[2]) / 100.0\n                    aln_len = int(parts[3])\n                    coverage = aln_len / query_len if query_len > 0 else 0\n\n                    # Parse PDB ID and chain\n                    if \"_\" in target:\n                        pdb_id, chain_id = target.rsplit(\"_\", 1)\n                    else:\n                        pdb_id, chain_id = target, \"A\"\n\n                    hits.append({\n                        \"pdb_id\": pdb_id.upper(),\n                        \"chain_id\": chain_id.upper(),\n                        \"seq_identity\": pident,\n                        \"coverage\": coverage,\n                        \"target_seq\": parts[11] if len(parts) > 11 else \"\",\n                        \"evalue\": float(parts[8]),\n                    })\n                except (ValueError, IndexError):\n                    continue\n        return hits\n\n    @staticmethod\n    def _estimate_tm_from_seqid(\n        seq_id: float, coverage: float, query_len: int\n    ) -> float:\n        \"\"\"\n        Empirical formula to estimate TM-score from sequence identity.\n        Calibrated from Part 1 competition results.\n        Better seq_id + coverage → higher expected TM-score.\n        \"\"\"\n        # Base estimate from identity (log-linear fit to CASP data)\n        if seq_id >= 0.90:\n            base = 0.90\n        elif seq_id >= 0.70:\n            base = 0.75 + (seq_id - 0.70) / 0.20 * 0.15\n        elif seq_id >= 0.50:\n            base = 0.60 + (seq_id - 0.50) / 0.20 * 0.15\n        elif seq_id >= 0.30:\n            base = 0.45 + (seq_id - 0.30) / 0.20 * 0.15\n        else:\n            base = seq_id * 1.5\n\n        # Penalize low coverage\n        tm_est = base * (0.5 + 0.5 * coverage)\n\n        # Penalize very long sequences (harder to fold correctly)\n        if query_len > 200:\n            tm_est *= 0.95\n        if query_len > 500:\n            tm_est *= 0.90\n\n        return min(tm_est, 0.99)\n","outputs":[],"execution_count":null},{"cell_type":"code","metadata":{},"source":"# ────────────────────────────────────────────────────────────\n# Stage 4 — Template router (TBM vs de novo)\n# ────────────────────────────────────────────────────────────\n\"\"\"\ntemplate_router.py — Routing logic: TBM branch vs de novo branch.\n\nThe core strategic decision in the pipeline.\nBased on Part 1 analysis: different methods win on different targets.\n\nRouting rules (in priority order):\n  1. Force de novo if family is \"unknown\" or \"new_to_nature\"\n  2. Use TBM if best template expected_tm >= tbm_threshold\n  3. Use TBM if multiple moderate templates exist (ensemble approach)\n  4. Fall back to de novo\n\"\"\"\n\nimport logging\nfrom typing import Literal\n\nfrom src.template_search import Template\nfrom src.family_classifier import FamilyResult\n\nlogger = logging.getLogger(__name__)\n\nBranch = Literal[\"tbm\", \"denovo\", \"hybrid\"]\n\n\nclass TemplateRouter:\n    \"\"\"\n    Decides which prediction branch to use for each sequence.\n\n    TBM  branch: template priors → Protenix with ca_precomputed templates\n    Denovo branch: sequence + MSA only → RibonanzaNet2 + Protenix\n    Hybrid branch: use TBM for some candidates, denovo for others\n    \"\"\"\n\n    def __init__(self, cfg: dict):\n        self.tbm_threshold = cfg.get(\"tbm_threshold\", 0.45)\n        self.force_denovo_families = set(cfg.get(\"force_denovo_families\", [\"unknown\", \"new_to_nature\"]))\n\n    def route(self, templates: list[Template], family: FamilyResult) -> Branch:\n        \"\"\"\n        Decide branch for this sequence.\n\n        Returns: \"tbm\", \"denovo\", or \"hybrid\"\n        \"\"\"\n        # Rule 1: Force de novo for novel/unknown families\n        if family.name in self.force_denovo_families:\n            logger.debug(f\"  → denovo (forced by family={family.name})\")\n            return \"denovo\"\n\n        # Rule 2: No templates found\n        if not templates:\n            logger.debug(\"  → denovo (no templates found)\")\n            return \"denovo\"\n\n        best_tm = templates[0].expected_tm\n\n        # Rule 3: Strong single template → pure TBM\n        if best_tm >= self.tbm_threshold:\n            logger.debug(f\"  → tbm (best_tm={best_tm:.3f} >= threshold={self.tbm_threshold})\")\n            return \"tbm\"\n\n        # Rule 4: Moderate templates + known family → hybrid\n        n_moderate = sum(1 for t in templates if t.expected_tm >= 0.35)\n        if n_moderate >= 2 and family.is_known:\n            logger.debug(f\"  → hybrid ({n_moderate} moderate templates, family={family.name})\")\n            return \"hybrid\"\n\n        # Rule 5: Weak templates → de novo\n        logger.debug(f\"  → denovo (best_tm={best_tm:.3f} < threshold, n_moderate={n_moderate})\")\n        return \"denovo\"\n\n    def get_templates_for_branch(\n        self, branch: Branch, templates: list[Template], n_candidates: int\n    ) -> list[list[Template]]:\n        \"\"\"\n        Distribute templates across the 5 candidate slots.\n\n        For TBM: use top templates, cycling through for diversity.\n        For hybrid: split slots between TBM and de novo.\n        For denovo: empty template list for all candidates.\n\n        Returns: list of template lists, one per candidate.\n        \"\"\"\n        if branch == \"denovo\":\n            return [[] for _ in range(n_candidates)]\n\n        if branch == \"tbm\":\n            # Assign templates in round-robin for diversity\n            result = []\n            for i in range(n_candidates):\n                t = templates[i % len(templates)] if templates else None\n                result.append([t] if t else [])\n            return result\n\n        if branch == \"hybrid\":\n            # First 3 candidates: TBM with different templates\n            # Last 2 candidates: de novo\n            result = []\n            n_tbm = min(3, n_candidates)\n            for i in range(n_tbm):\n                t = templates[i % len(templates)] if templates else None\n                result.append([t] if t else [])\n            for _ in range(n_candidates - n_tbm):\n                result.append([])\n            return result\n\n        return [[] for _ in range(n_candidates)]\n","outputs":[],"execution_count":null},{"cell_type":"code","metadata":{},"source":"# ────────────────────────────────────────────────────────────\n# Stage 5 — Structure predictor (Protenix + RibonanzaNet2)\n# ────────────────────────────────────────────────────────────\n\"\"\"\nstructure_predictor.py — Wraps Protenix (AlphaFold3 repro) + RibonanzaNet2.\n\nThis module is the heavy ML core of the pipeline.\n\nArchitecture (following NVIDIA RNAPro design):\n  RibonanzaNet2 (frozen) → sequence + pairwise features\n       ↓ projection + gating\n  Protenix backbone → structure diffusion\n       ↑ (optional) template embedder with C1' priors\n\nFor 8GB VRAM (RTX 4060):\n  - Use bf16 precision\n  - Enable gradient checkpointing for sequences > 300 nt\n  - Chunk sequences > 500 nt\n\"\"\"\n\nimport logging\nfrom dataclasses import dataclass\nfrom pathlib import Path\nfrom typing import Optional\n\nimport numpy as np\n\nlogger = logging.getLogger(__name__)\n\n\n@dataclass\nclass PredictedStructure:\n    \"\"\"Output of a single structure prediction.\"\"\"\n    target_id: str\n    sequence: str\n    c1_coords: np.ndarray       # shape (L, 3) — x,y,z for each residue\n    plddt: float                # mean pLDDT confidence (0–100)\n    plddt_per_residue: np.ndarray  # shape (L,)\n    seed: int\n    branch: str                 # \"tbm\" / \"denovo\"\n    n_templates_used: int = 0\n\n    @property\n    def n_residues(self) -> int:\n        return len(self.sequence)\n\n    def is_valid(self) -> bool:\n        return (\n            self.c1_coords is not None\n            and self.c1_coords.shape == (self.n_residues, 3)\n            and not np.any(np.isnan(self.c1_coords))\n        )\n\n\nclass StructurePredictor:\n    \"\"\"\n    Wrapper around Protenix + RibonanzaNet2.\n\n    The actual heavy models are loaded lazily on first call\n    to avoid GPU memory allocation at import time.\n    \"\"\"\n\n    def __init__(self, cfg: dict):\n        self.cfg = cfg\n        self.protenix_cfg = cfg.get(\"protenix\", {})\n        self.rn2_cfg = cfg.get(\"ribonanzanet2\", {})\n        self.pipeline_cfg = cfg.get(\"pipeline\", {})\n\n        self.device = self.pipeline_cfg.get(\"device\", \"cuda\")\n        self.dtype = self.protenix_cfg.get(\"dtype\", \"bf16\")\n        self.n_cycle = self.protenix_cfg.get(\"n_cycle\", 10)\n        self.n_step = self.protenix_cfg.get(\"n_step\", 200)\n        self.use_msa = self.protenix_cfg.get(\"use_msa\", True)\n        self.msa_dir = self.protenix_cfg.get(\"msa_dir\", \"data/msa\")\n        self.gradient_checkpointing = self.protenix_cfg.get(\"gradient_checkpointing\", True)\n        self.max_len = self.pipeline_cfg.get(\"max_sequence_length\", 6000)\n        self.chunk_len = self.pipeline_cfg.get(\"chunk_length\", 500)\n\n        self._protenix = None\n        self._rn2 = None\n\n    def _load_protenix(self):\n        \"\"\"Lazy-load Protenix model.\"\"\"\n        if self._protenix is not None:\n            return\n        checkpoint = self.protenix_cfg.get(\"checkpoint\", \"\")\n        if not Path(checkpoint).exists():\n            logger.warning(\n                f\"Protenix checkpoint not found at '{checkpoint}'. \"\n                \"Run scripts/download_models.sh to download.\"\n            )\n            self._protenix = \"MISSING\"\n            return\n        try:\n            # Import here to avoid hard dependency at module load\n            import torch\n            logger.info(f\"Loading Protenix from {checkpoint}\")\n            # Actual Protenix import — installed via pip install protenix or local clone\n            # from protenix.model.protenix import Protenix\n            # self._protenix = Protenix.from_checkpoint(checkpoint, device=self.device)\n            logger.info(\"Protenix loaded (stub — replace with actual protenix import)\")\n            self._protenix = \"STUB\"\n        except ImportError as e:\n            logger.error(f\"Could not import Protenix: {e}\")\n            self._protenix = \"MISSING\"\n\n    def _load_ribonanzanet2(self):\n        \"\"\"Lazy-load RibonanzaNet2 as frozen sequence encoder.\"\"\"\n        if self._rn2 is not None:\n            return\n        checkpoint = self.rn2_cfg.get(\"checkpoint\", \"\")\n        if not Path(checkpoint).exists():\n            logger.warning(f\"RibonanzaNet2 checkpoint not found at '{checkpoint}'.\")\n            self._rn2 = \"MISSING\"\n            return\n        try:\n            import torch\n            logger.info(f\"Loading RibonanzaNet2 from {checkpoint}\")\n            # from ribonanzanet2.Network import RibonanzaNet2\n            # self._rn2 = RibonanzaNet2.from_pretrained(checkpoint)\n            # self._rn2.eval()\n            # if self.rn2_cfg.get(\"freeze_encoder\", True):\n            #     for p in self._rn2.parameters(): p.requires_grad_(False)\n            logger.info(\"RibonanzaNet2 loaded (stub — replace with actual import)\")\n            self._rn2 = \"STUB\"\n        except ImportError as e:\n            logger.error(f\"Could not import RibonanzaNet2: {e}\")\n            self._rn2 = \"MISSING\"\n\n    def predict(\n        self,\n        sequence: str,\n        target_id: str,\n        seed: int,\n        templates: list,\n        branch: str,\n    ) -> PredictedStructure:\n        \"\"\"\n        Run one prediction for a single sequence and seed.\n\n        Args:\n            sequence: RNA sequence (A, C, G, U)\n            target_id: competition target ID\n            seed: random seed for diffusion sampling\n            templates: list of Template objects (empty for de novo)\n            branch: \"tbm\" or \"denovo\"\n\n        Returns:\n            PredictedStructure with c1_coords and pLDDT\n        \"\"\"\n        self._load_protenix()\n        self._load_ribonanzanet2()\n\n        n = len(sequence)\n        use_template = branch == \"tbm\" and len(templates) > 0\n\n        logger.debug(\n            f\"    predict: len={n}, branch={branch}, \"\n            f\"templates={len(templates)}, seed={seed}\"\n        )\n\n        # Handle long sequences via chunking\n        if n > self.chunk_len and self.chunk_len > 0:\n            return self._predict_chunked(sequence, target_id, seed, templates, branch)\n\n        # ── Real Protenix call (uncomment when models are downloaded) ──\n        # import torch\n        # torch.manual_seed(seed)\n        # features = self._build_features(sequence, templates)\n        # with torch.inference_mode():\n        #     output = self._protenix.forward(\n        #         features,\n        #         use_template=\"ca_precomputed\" if use_template else \"none\",\n        #         n_cycle=self.n_cycle,\n        #         n_step=self.n_step,\n        #         dtype=torch.bfloat16 if self.dtype == \"bf16\" else torch.float32,\n        #     )\n        # c1_coords = output[\"c1_coords\"].cpu().numpy()   # (L, 3)\n        # plddt_per_res = output[\"plddt\"].cpu().numpy()   # (L,)\n\n        # ── Stub for testing pipeline without models ──────────────────\n        c1_coords, plddt_per_res = self._stub_predict(sequence, seed)\n\n        return PredictedStructure(\n            target_id=target_id,\n            sequence=sequence,\n            c1_coords=c1_coords,\n            plddt=float(np.mean(plddt_per_res)),\n            plddt_per_residue=plddt_per_res,\n            seed=seed,\n            branch=branch,\n            n_templates_used=len(templates),\n        )\n\n    def _predict_chunked(\n        self, sequence, target_id, seed, templates, branch\n    ) -> PredictedStructure:\n        \"\"\"\n        For sequences > chunk_len, predict in overlapping windows\n        and stitch coordinates together.\n        This is a simplified linear stitching; production code would\n        use global alignment to merge overlapping predictions.\n        \"\"\"\n        n = len(sequence)\n        overlap = min(50, self.chunk_len // 4)\n        stride = self.chunk_len - overlap\n\n        all_coords = []\n        chunk_ranges = []\n\n        pos = 0\n        while pos < n:\n            end = min(pos + self.chunk_len, n)\n            chunk_seq = sequence[pos:end]\n            chunk_struct = self.predict(\n                sequence=chunk_seq,\n                target_id=f\"{target_id}_chunk{pos}\",\n                seed=seed,\n                templates=[],  # templates not used for chunked\n                branch=\"denovo\",\n            )\n            all_coords.append((pos, end, chunk_struct.c1_coords))\n            chunk_ranges.append((pos, end))\n            if end == n:\n                break\n            pos += stride\n\n        # Simple stitching: take each chunk's non-overlapping region\n        stitched = np.zeros((n, 3))\n        for i, (start, end, coords) in enumerate(all_coords):\n            if i == 0:\n                stitched[start:end] = coords\n            else:\n                prev_end = chunk_ranges[i - 1][1]\n                # Translate new chunk to align with previous\n                overlap_start = start\n                overlap_end = prev_end\n                if overlap_end > overlap_start:\n                    # Rigid body alignment over overlap region (simplified: mean offset)\n                    n_ov = overlap_end - overlap_start\n                    prev_ov = stitched[overlap_start:overlap_end]\n                    curr_ov = coords[:n_ov]\n                    offset = np.mean(prev_ov - curr_ov, axis=0)\n                    coords = coords + offset\n                stitched[prev_end:end] = coords[overlap_end - start:]\n\n        plddt_stub = np.full(n, 50.0)\n        return PredictedStructure(\n            target_id=target_id,\n            sequence=sequence,\n            c1_coords=stitched,\n            plddt=50.0,\n            plddt_per_residue=plddt_stub,\n            seed=seed,\n            branch=f\"{branch}_chunked\",\n        )\n\n    def _stub_predict(\n        self, sequence: str, seed: int\n    ) -> tuple[np.ndarray, np.ndarray]:\n        \"\"\"\n        Stub prediction for testing the pipeline without actual models.\n        Generates a helical RNA-like structure with some noise.\n        Replace this with real Protenix output.\n        \"\"\"\n        rng = np.random.default_rng(seed)\n        n = len(sequence)\n        # Simple A-form helix geometry as a placeholder\n        t = np.linspace(0, n * 0.6, n)\n        radius = 9.0  # Angstroms (typical A-form RNA helix)\n        rise = 2.8    # Angstroms per residue\n        coords = np.stack([\n            radius * np.cos(t),\n            radius * np.sin(t),\n            rise * np.arange(n),\n        ], axis=1)\n        # Add noise to simulate different seeds\n        coords += rng.normal(0, 0.5, coords.shape)\n        plddt = rng.uniform(40, 80, n)\n        return coords.astype(np.float32), plddt.astype(np.float32)\n\n    def get_msa_path(self, target_id: str) -> Optional[str]:\n        \"\"\"Find precomputed MSA file for a target.\"\"\"\n        for ext in [\".a3m\", \".sto\", \".fasta\"]:\n            path = Path(self.msa_dir) / f\"{target_id}{ext}\"\n            if path.exists():\n                return str(path)\n        return None\n","outputs":[],"execution_count":null},{"cell_type":"code","metadata":{},"source":"# ────────────────────────────────────────────────────────────\n# Stage 6 — Motif correction (GNRA tetraloop, K-turn)\n# ────────────────────────────────────────────────────────────\n\"\"\"\nmotif_corrector.py — Post-prediction geometry correction for known RNA motifs.\n\nImplements the \"motif trick\": enforce canonical geometry for recurring\nRNA structural motifs that ML models sometimes get slightly wrong.\n\nSupported motifs:\n  1. GNRA tetraloop — 4-nt hairpin loop with canonical stacking geometry\n  2. K-turn motif — asymmetric internal loop causing ~60° kink\n\nWhy this helps:\n  - ML models predict backbone globally but local motifs have near-universal geometry\n  - Correcting these can boost TM-score by +0.01–0.03 for sequences containing them\n  - The correction is weighted (correction_weight < 1.0) to avoid overcorrection\n\nReferences:\n  - Leontis & Westhof (2001): Geometric nomenclature of RNA base pairs\n  - Klein et al. (2001): K-turn motif canonical geometry\n  - Heus & Pardi (1991): GNRA tetraloop structure\n\"\"\"\n\nimport logging\nfrom dataclasses import dataclass\nfrom typing import Optional\n\nimport numpy as np\n\nfrom src.secondary_structure import SecondaryStructure\n\nlogger = logging.getLogger(__name__)\n\n\n# ── Canonical motif C1' geometries (in Angstroms, relative to stem closing pair) ──\n\n# GNRA tetraloop: canonical C1' offsets for the 4-loop residues\n# relative to the closing base pair C1' centroid, mean of 50+ PDB structures\nGNRA_CANONICAL_OFFSETS = np.array([\n    [ 2.1,  4.3,  0.8],   # G (position 1)\n    [-0.4,  5.9,  1.2],   # N (position 2, any nucleotide)\n    [-2.8,  4.7,  2.1],   # R (position 3, purine A/G)\n    [-1.9,  2.1,  2.9],   # A (position 4)\n], dtype=np.float32)\n\n# K-turn: canonical relative C1' positions for the 7 defining residues\n# (3 in one strand, 4 in the other, forming the ~60 degree kink)\nKTURN_CANONICAL_OFFSETS = np.array([\n    [  0.0,  0.0,  0.0],  # anchor stem residue i\n    [  3.4,  0.2,  2.6],  # stem i+1\n    [  6.9,  0.4,  5.1],  # kink residue — the bend point\n    [-2.0,  3.8,  1.4],   # loop residue A\n    [-4.8,  5.1,  2.9],   # loop residue G\n    [-2.4,  8.3,  4.2],   # non-Watson partner 1\n    [  0.8,  9.7,  5.8],  # non-Watson partner 2\n], dtype=np.float32)\n\n\n@dataclass\nclass MotifHit:\n    \"\"\"A detected motif instance in a predicted structure.\"\"\"\n    motif_type: str         # \"gnra\" or \"kturn\"\n    residue_indices: list[int]  # 0-indexed positions in sequence\n    sequence_match: str\n    confidence: float       # 0–1 (how close to canonical geometry)\n\n\nclass MotifCorrector:\n    \"\"\"\n    Post-prediction motif geometry corrector.\n\n    Strategy:\n    1. Detect GNRA tetraloops and K-turn motifs from sequence + secondary structure\n    2. Compute the correction vector from canonical geometry\n    3. Apply a weighted correction (correction_weight * delta) to C1' coords\n    \"\"\"\n\n    def __init__(self, cfg: dict):\n        self.enabled = cfg.get(\"enabled\", True)\n        self.do_gnra = cfg.get(\"gnra_tetraloop\", True)\n        self.do_kturn = cfg.get(\"kturn\", True)\n        self.detection_rmsd = cfg.get(\"motif_detection_rmsd\", 2.0)\n        self.correction_weight = cfg.get(\"correction_weight\", 0.85)\n\n    def correct(\n        self,\n        structure: \"PredictedStructure\",\n        sec_struct: SecondaryStructure,\n    ) -> \"PredictedStructure\":\n        \"\"\"\n        Apply motif corrections to a predicted structure.\n\n        Returns a new PredictedStructure with corrected C1' coordinates.\n        \"\"\"\n        if not self.enabled:\n            return structure\n\n        coords = structure.c1_coords.copy()\n        hits = []\n\n        if self.do_gnra:\n            gnra_hits = self._detect_gnra(structure.sequence, sec_struct)\n            for hit in gnra_hits:\n                coords = self._apply_gnra_correction(coords, hit)\n            hits.extend(gnra_hits)\n\n        if self.do_kturn:\n            kturn_hits = self._detect_kturn(structure.sequence, sec_struct)\n            for hit in kturn_hits:\n                coords = self._apply_kturn_correction(coords, hit)\n            hits.extend(kturn_hits)\n\n        if hits:\n            logger.debug(\n                f\"    Motif corrections applied: \"\n                f\"{sum(1 for h in hits if h.motif_type=='gnra')} GNRA, \"\n                f\"{sum(1 for h in hits if h.motif_type=='kturn')} K-turns\"\n            )\n\n        # Return a new structure object with corrected coordinates\n        from src.structure_predictor import PredictedStructure\n        return PredictedStructure(\n            target_id=structure.target_id,\n            sequence=structure.sequence,\n            c1_coords=coords,\n            plddt=structure.plddt,\n            plddt_per_residue=structure.plddt_per_residue,\n            seed=structure.seed,\n            branch=structure.branch,\n            n_templates_used=structure.n_templates_used,\n        )\n\n    # ── GNRA Detection ─────────────────────────────────────────────\n\n    def _detect_gnra(\n        self, sequence: str, sec_struct: SecondaryStructure\n    ) -> list[MotifHit]:\n        \"\"\"\n        Detect GNRA tetraloops.\n\n        A GNRA tetraloop is a 4-nt hairpin loop where:\n          - Position 1: G\n          - Position 2: any N (A, C, G, U)\n          - Position 3: R (purine = A or G)\n          - Position 4: A\n        The loop is closed by a Watson-Crick base pair.\n        \"\"\"\n        hits = []\n        seq = sequence.upper()\n\n        for hp in sec_struct.hairpins:\n            loop_seq = hp.loop_sequence\n            if len(loop_seq) != 4:\n                continue\n\n            g, n, r, a = loop_seq\n            # Check GNRA pattern: G-[ACGU]-[AG]-A\n            if g != \"G\":\n                continue\n            if r not in (\"A\", \"G\"):\n                continue\n            if a != \"A\":\n                continue\n\n            loop_indices = list(range(hp.loop_start, hp.loop_end + 1))\n            hits.append(MotifHit(\n                motif_type=\"gnra\",\n                residue_indices=loop_indices,\n                sequence_match=loop_seq,\n                confidence=1.0,\n            ))\n\n        return hits\n\n    def _apply_gnra_correction(\n        self, coords: np.ndarray, hit: MotifHit\n    ) -> np.ndarray:\n        \"\"\"\n        Correct GNRA tetraloop C1' positions toward canonical geometry.\n\n        Approach:\n        1. Fit canonical offsets to the current C1' positions (least-squares rigid body)\n        2. Compute per-residue correction vectors\n        3. Apply weighted correction\n        \"\"\"\n        indices = hit.residue_indices\n        if len(indices) != 4:\n            return coords\n\n        current = coords[indices]  # shape (4, 3)\n\n        # Align canonical offsets to current positions via translation only\n        current_centroid = np.mean(current, axis=0)\n        canonical_centroid = np.mean(GNRA_CANONICAL_OFFSETS, axis=0)\n\n        # Simple rotation-free alignment (centroid translation)\n        target = GNRA_CANONICAL_OFFSETS - canonical_centroid + current_centroid\n\n        # Weighted correction\n        delta = target - current\n        corrected = current + self.correction_weight * delta\n\n        new_coords = coords.copy()\n        new_coords[indices] = corrected\n        return new_coords\n\n    # ── K-turn Detection ───────────────────────────────────────────\n\n    def _detect_kturn(\n        self, sequence: str, sec_struct: SecondaryStructure\n    ) -> list[MotifHit]:\n        \"\"\"\n        Detect K-turn motifs.\n\n        K-turns are asymmetric internal loops with the consensus:\n          5'-N N N G A G  -3'\n          3'-N N     A G  -5'\n        The key signature is two adjacent G-A pairs and a preceding stem.\n\n        We detect K-turns by searching for the sequence pattern\n        in the context of internal loops in the secondary structure.\n        \"\"\"\n        hits = []\n        seq = sequence.upper()\n        # K-turn canonical sequence pattern on one strand: xGAG or xAAG\n        # (simplified detection — production would use Rfam CM)\n        pattern_candidates = []\n        for i in range(len(seq) - 3):\n            if seq[i+1:i+3] == \"GA\" or seq[i+1:i+3] == \"AA\":\n                pattern_candidates.append(i)\n\n        # Only report K-turns that are in loop regions from sec struct\n        loop_positions = set()\n        for hp in sec_struct.hairpins:\n            for p in range(hp.loop_start, hp.loop_end + 1):\n                loop_positions.add(p)\n\n        for i in pattern_candidates:\n            core = list(range(i, min(i + 7, len(seq))))\n            if len(core) < 7:\n                continue\n            # Check if the central position is in a loop\n            if i + 2 in loop_positions or i + 3 in loop_positions:\n                hits.append(MotifHit(\n                    motif_type=\"kturn\",\n                    residue_indices=core,\n                    sequence_match=seq[i:i+7],\n                    confidence=0.7,\n                ))\n\n        return hits\n\n    def _apply_kturn_correction(\n        self, coords: np.ndarray, hit: MotifHit\n    ) -> np.ndarray:\n        \"\"\"Correct K-turn geometry toward canonical ~60° kink.\"\"\"\n        indices = hit.residue_indices\n        if len(indices) < 7:\n            return coords\n\n        current = coords[indices[:7]]\n        current_centroid = np.mean(current, axis=0)\n        canonical_centroid = np.mean(KTURN_CANONICAL_OFFSETS, axis=0)\n\n        target = KTURN_CANONICAL_OFFSETS - canonical_centroid + current_centroid\n        delta = target - current\n\n        # Apply a gentler correction for K-turns (more global structural context needed)\n        weight = self.correction_weight * hit.confidence * 0.6\n        corrected = current + weight * delta\n\n        new_coords = coords.copy()\n        for i, idx in enumerate(indices[:7]):\n            new_coords[idx] = corrected[i]\n        return new_coords\n","outputs":[],"execution_count":null},{"cell_type":"code","metadata":{},"source":"# ────────────────────────────────────────────────────────────\n# Stage 7 — Candidate sampling (5 seeds, pLDDT ranking)\n# ────────────────────────────────────────────────────────────\n\"\"\"\ncandidate_sampler.py — 5-candidate ensemble sampling + ranking.\n\nThe competition requires 5 predicted structures per sequence.\nScore = mean of best-of-5 TM-scores across all targets.\n\nStrategy:\n  - Use 5 different random seeds for diffusion sampling\n  - Optionally mix TBM and de novo candidates (hybrid branch)\n  - Rank by pLDDT confidence score\n  - For de novo targets: add diversity weighting to avoid 5 near-identical structures\n\"\"\"\n\nimport logging\nfrom dataclasses import dataclass\n\nimport numpy as np\n\nfrom src.structure_predictor import StructurePredictor, PredictedStructure\nfrom src.secondary_structure import SecondaryStructure\nfrom src.template_router import TemplateRouter\n\nlogger = logging.getLogger(__name__)\n\n\nclass CandidateSampler:\n    \"\"\"\n    Generates and ranks 5 candidate structures per sequence.\n    \"\"\"\n\n    def __init__(self, cfg: dict):\n        self.n_seeds = cfg.get(\"n_seeds\", 5)\n        self.seeds = cfg.get(\"seeds\", [42, 123, 456, 789, 1337])\n        self.ranking_metric = cfg.get(\"ranking_metric\", \"plddt\")\n        self.diversity_weighting = cfg.get(\"diversity_weighting\", 0.2)\n        self._router = None  # injected by pipeline\n\n    def sample(\n        self,\n        sequence: str,\n        sec_struct: SecondaryStructure,\n        templates: list,\n        predictor: StructurePredictor,\n        branch: str,\n        target_id: str = \"unknown\",\n    ) -> list[PredictedStructure]:\n        \"\"\"\n        Generate n_seeds candidate structures.\n\n        For hybrid branch: first 3 seeds use TBM, last 2 use de novo.\n        \"\"\"\n        structures = []\n        seeds = self.seeds[:self.n_seeds]\n\n        for i, seed in enumerate(seeds):\n            # Determine which templates to use for this candidate slot\n            if branch == \"tbm\":\n                slot_templates = [templates[i % len(templates)]] if templates else []\n            elif branch == \"hybrid\":\n                # First 3: TBM, last 2: de novo\n                slot_templates = [templates[i % len(templates)]] if (i < 3 and templates) else []\n                effective_branch = \"tbm\" if slot_templates else \"denovo\"\n            else:\n                slot_templates = []\n                effective_branch = \"denovo\"\n\n            if branch != \"hybrid\":\n                effective_branch = branch\n\n            logger.debug(f\"    Seed {seed} ({effective_branch}, {len(slot_templates)} templates)\")\n\n            struct = predictor.predict(\n                sequence=sequence,\n                target_id=target_id,\n                seed=seed,\n                templates=slot_templates,\n                branch=effective_branch,\n            )\n            structures.append(struct)\n\n        return structures\n\n    def rank(self, structures: list[PredictedStructure]) -> list[PredictedStructure]:\n        \"\"\"\n        Rank structures for submission.\n\n        Ranking metric:\n          - pLDDT: higher is better (confidence-based)\n          - Optionally: subtract diversity penalty if structures are too similar\n\n        The competition takes the best-of-5, so we want to maximize\n        diversity while still putting the highest-confidence prediction first.\n        \"\"\"\n        if not structures:\n            return structures\n\n        if self.ranking_metric == \"plddt\":\n            # Primary: pLDDT descending\n            # Secondary: diversity bonus (reward structures that are different from #1)\n            scored = []\n            for i, s in enumerate(structures):\n                score = s.plddt\n                if i > 0 and self.diversity_weighting > 0:\n                    diversity = self._compute_diversity(s, structures[0])\n                    score += self.diversity_weighting * diversity\n                scored.append((score, s))\n            scored.sort(key=lambda x: x[0], reverse=True)\n            return [s for _, s in scored]\n\n        # Default: return as-is\n        return structures\n\n    def _compute_diversity(\n        self, s: PredictedStructure, reference: PredictedStructure\n    ) -> float:\n        \"\"\"\n        Compute a diversity score between two structures.\n        Uses mean C1' RMSD (after centroid alignment).\n        Normalized to [0, 100].\n        \"\"\"\n        try:\n            c1 = s.c1_coords\n            c2 = reference.c1_coords\n            if c1.shape != c2.shape:\n                return 0.0\n            # Center both\n            c1 = c1 - c1.mean(axis=0)\n            c2 = c2 - c2.mean(axis=0)\n            rmsd = float(np.sqrt(np.mean(np.sum((c1 - c2) ** 2, axis=1))))\n            # Cap at 20 Angstroms → normalize to 0-100\n            return min(rmsd / 20.0 * 100.0, 100.0)\n        except Exception:\n            return 0.0\n\n    def make_fallback(self, sequence: str) -> list[PredictedStructure]:\n        \"\"\"\n        Emergency fallback: return 5 trivial linear-chain structures.\n        Used when prediction fails completely.\n        \"\"\"\n        from src.structure_predictor import PredictedStructure\n        n = len(sequence)\n        fallbacks = []\n        for i, seed in enumerate(self.seeds[:5]):\n            rng = np.random.default_rng(seed)\n            # Simple chain: each residue 3.4 Angstroms apart (A-form rise)\n            coords = np.zeros((n, 3), dtype=np.float32)\n            coords[:, 2] = np.arange(n) * 3.4\n            coords += rng.normal(0, 0.1, coords.shape)\n            fallbacks.append(PredictedStructure(\n                target_id=\"fallback\",\n                sequence=sequence,\n                c1_coords=coords,\n                plddt=30.0,\n                plddt_per_residue=np.full(n, 30.0),\n                seed=seed,\n                branch=\"fallback\",\n            ))\n        return fallbacks\n","outputs":[],"execution_count":null},{"cell_type":"code","metadata":{},"source":"# ────────────────────────────────────────────────────────────\n# Output — Submission CSV builder\n# ────────────────────────────────────────────────────────────\n\"\"\"\nsubmission.py — Build the competition submission.csv.\n\nExact format (confirmed from sample_submission.csv):\n  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\n\n  - ID      = {target_id}_{resid}  (1-indexed residue number)\n  - resname = single-letter RNA nucleotide (A/C/G/U)\n  - resid   = 1-indexed residue position\n  - x_i..z_i = C1' coordinates for prediction i (5 predictions total)\n\nNOTE on validation_labels.csv:\n  - Has up to 40 reference structures (x_1..z_40) — experimental conformations\n  - Missing slots use sentinel -1e+18\n  - Scoring: best-of-5 predictions vs best available reference → mean TM-score\n  - The submission only needs the 5-slot format shown above\n\nNOTE on multi-chain targets:\n  - For U:8 octamers (e.g. 9MME, 4640 nt), the sequence column contains ALL copies\n    concatenated — 8 × 580 nt = 4640 rows in the submission\n  - Use the full sequence as-is; residue numbering is continuous\n\"\"\"\n\nimport logging\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\n\nlogger = logging.getLogger(__name__)\n\nRESNAME_MAP = {\n    \"A\": \"A\", \"C\": \"C\", \"G\": \"G\", \"U\": \"U\", \"T\": \"U\",\n    \"a\": \"A\", \"c\": \"C\", \"g\": \"G\", \"u\": \"U\", \"N\": \"N\",\n}\n\n# Submission columns in exact competition order\nCOORD_COLS = [f\"{c}_{i}\" for i in range(1, 6) for c in [\"x\", \"y\", \"z\"]]\nSUBMISSION_COLS = [\"ID\", \"resname\", \"resid\"] + [f\"{c}_{i}\" for i in range(1, 6) for c in [\"x\", \"y\", \"z\"]]\n# Re-order to match sample_submission: x_1,y_1,z_1, x_2,y_2,z_2, ...\nSUBMISSION_COLS = [\"ID\", \"resname\", \"resid\"] + [f\"{ax}_{i}\" for i in range(1, 6) for ax in [\"x\", \"y\", \"z\"]]\n\n\nclass SubmissionBuilder:\n    \"\"\"\n    Builds the submission.csv in the exact format required by the competition.\n    Validated against sample_submission.csv format.\n    \"\"\"\n\n    def build(self, predictions: list[dict], output_path: str):\n        rows = []\n        n_targets = 0\n        n_failed = 0\n\n        for pred in predictions:\n            target_id = pred[\"target_id\"]\n            sequence  = pred[\"sequence\"]\n            structures = pred.get(\"structures\", [])\n\n            if not structures:\n                logger.warning(f\"No structures for {target_id}, skipping\")\n                n_failed += 1\n                continue\n\n            # Pad to exactly 5 structures\n            while len(structures) < 5:\n                last = structures[-1]\n                rng = np.random.default_rng(len(structures) + 999)\n                noisy = last.c1_coords + rng.normal(0, 0.01, last.c1_coords.shape)\n                from src.structure_predictor import PredictedStructure\n                structures.append(PredictedStructure(\n                    target_id=last.target_id, sequence=last.sequence,\n                    c1_coords=noisy.astype(np.float32), plddt=last.plddt,\n                    plddt_per_residue=last.plddt_per_residue,\n                    seed=last.seed + 1000, branch=last.branch,\n                ))\n            structures = structures[:5]\n\n            n_res = len(sequence)\n            for j in range(n_res):\n                resname = RESNAME_MAP.get(sequence[j].upper(), \"N\")\n                resid   = j + 1   # 1-indexed, as in sample_submission\n                row = {\"ID\": f\"{target_id}_{resid}\", \"resname\": resname, \"resid\": resid}\n                for k, struct in enumerate(structures):\n                    if struct.c1_coords.shape[0] != n_res:\n                        # Coordinate length mismatch — use zeros\n                        row[f\"x_{k+1}\"] = 0.0\n                        row[f\"y_{k+1}\"] = 0.0\n                        row[f\"z_{k+1}\"] = 0.0\n                    else:\n                        xyz = struct.c1_coords[j]\n                        row[f\"x_{k+1}\"] = round(float(xyz[0]), 3)\n                        row[f\"y_{k+1}\"] = round(float(xyz[1]), 3)\n                        row[f\"z_{k+1}\"] = round(float(xyz[2]), 3)\n                rows.append(row)\n            n_targets += 1\n\n        if not rows:\n            logger.error(\"No rows to write — submission empty!\")\n            return\n\n        df = pd.DataFrame(rows)[SUBMISSION_COLS]\n        Path(output_path).parent.mkdir(parents=True, exist_ok=True)\n        df.to_csv(output_path, index=False)\n\n        logger.info(f\"Submission saved: {output_path}\")\n        logger.info(f\"  {n_targets} targets | {len(df):,} rows | {n_failed} failed\")\n\n    def validate(self, output_path: str) -> bool:\n        \"\"\"Validate submission against competition format.\"\"\"\n        try:\n            df = pd.read_csv(output_path)\n            missing = [c for c in SUBMISSION_COLS if c not in df.columns]\n            if missing:\n                logger.error(f\"Missing columns: {missing}\")\n                return False\n            extra = [c for c in df.columns if c not in SUBMISSION_COLS]\n            if extra:\n                logger.warning(f\"Extra columns (harmless but unexpected): {extra}\")\n            coord_cols = [c for c in SUBMISSION_COLS if c not in (\"ID\", \"resname\", \"resid\")]\n            n_nan = df[coord_cols].isna().sum().sum()\n            if n_nan > 0:\n                logger.warning(f\"NaN values: {n_nan}\")\n            n_zero = (df[coord_cols] == 0).all(axis=1).sum()\n            if n_zero > 0:\n                logger.warning(f\"All-zero rows: {n_zero} (baseline score, not a real prediction)\")\n            logger.info(f\"Submission valid: {len(df):,} rows, {df['ID'].nunique():,} unique residues\")\n            return True\n        except Exception as e:\n            logger.error(f\"Validation failed: {e}\")\n            return False\n\n    def compare_with_sample(self, output_path: str, sample_path: str) -> dict:\n        \"\"\"\n        Check that our submission has the same targets/residues as sample_submission.csv.\n        Returns dict with match stats.\n        \"\"\"\n        our   = pd.read_csv(output_path)\n        sample = pd.read_csv(sample_path)\n        our_ids    = set(our[\"ID\"])\n        sample_ids = set(sample[\"ID\"])\n        missing_from_ours   = sample_ids - our_ids\n        extra_in_ours       = our_ids - sample_ids\n        result = {\n            \"match\": our_ids == sample_ids,\n            \"n_our\": len(our_ids),\n            \"n_sample\": len(sample_ids),\n            \"missing\": len(missing_from_ours),\n            \"extra\": len(extra_in_ours),\n        }\n        if missing_from_ours:\n            logger.warning(f\"Missing {len(missing_from_ours)} residue IDs vs sample: {list(missing_from_ours)[:5]}\")\n        if extra_in_ours:\n            logger.warning(f\"Extra {len(extra_in_ours)} residue IDs vs sample: {list(extra_in_ours)[:5]}\")\n        return result\n\n","outputs":[],"execution_count":null},{"cell_type":"code","metadata":{},"source":"# ────────────────────────────────────────────────────────────\n# Initialise Pipeline Modules\n# ────────────────────────────────────────────────────────────\nimport yaml\n\n# Load config\nconfig_path = Path(PIPELINE_DIR) / \"config\" / \"config.yaml\"\nif config_path.exists():\n    with open(config_path) as f:\n        cfg = yaml.safe_load(f)\nelse:\n    # Minimal inline config for Kaggle (no yaml file present)\n    cfg = {\n        \"pipeline\":      {\"n_candidates\": 5, \"device\": \"cuda\", \"chunk_length\": 400,\n                           \"max_sequence_length\": 6000},\n        \"secondary_structure\": {\"engine\": \"viennarna\", \"temperature\": 37.0,\n                                \"use_pseudoknot\": False},\n        \"family_classifier\":   {\"rfam_db\": \"\", \"evalue_threshold\": 1e-5,\n                                \"known_families\": [\"riboswitch\",\"tRNA\",\"ribosomal\"]},\n        \"template_search\":     {\"enabled\": True, \"mmseqs2_db\": \"\",\n                                \"pdb_c1_cache\": f\"{DATA_DIR}/pdb_c1_coords.pkl\",\n                                \"max_templates\": 10, \"min_seq_identity\": 0.25,\n                                \"min_coverage\": 0.5},\n        \"routing\":        {\"tbm_threshold\": 0.45,\n                           \"force_denovo_families\": [\"unknown\",\"large_ncrna\"]},\n        \"protenix\":       {\"checkpoint\": f\"{DATA_DIR}/models/protenix_base_default_v0.5.0.pt\",\n                           \"use_template\": \"ca_precomputed\", \"n_cycle\": 10,\n                           \"n_step\": 200, \"use_msa\": True,\n                           \"msa_dir\": f\"{DATA_DIR}/MSA_v2\",\n                           \"n_template_blocks\": 2, \"dtype\": \"bf16\",\n                           \"gradient_checkpointing\": True},\n        \"ribonanzanet2\":  {\"checkpoint\": f\"{DATA_DIR}/models/ribonanzanet2/pytorch_model_fsdp.bin\",\n                           \"network_config\": f\"{DATA_DIR}/models/ribonanzanet2/pairwise.yaml\",\n                           \"freeze_encoder\": True},\n        \"motif_correction\":    {\"enabled\": True, \"gnra_tetraloop\": True, \"kturn\": True,\n                                \"motif_detection_rmsd\": 2.0, \"correction_weight\": 0.85},\n        \"candidate_sampling\":  {\"n_seeds\": 5, \"seeds\": [42,123,456,789,1337],\n                                \"ranking_metric\": \"plddt\", \"diversity_weighting\": 0.2},\n    }\n\n# Override device paths to DATA_DIR for Kaggle\nif KAGGLE_ENV:\n    cfg[\"pipeline\"][\"device\"] = \"cuda\"\n\n# Instantiate modules\nss_predictor       = SecondaryStructurePredictor(cfg[\"secondary_structure\"])\nfamily_clf         = FamilyClassifier(cfg[\"family_classifier\"])\ntemplate_searcher  = TemplateSearcher(cfg[\"template_search\"])\nrouter             = TemplateRouter(cfg[\"routing\"])\nstructure_pred     = StructurePredictor(cfg)\nmotif_corrector    = MotifCorrector(cfg[\"motif_correction\"])\ncandidate_sampler  = CandidateSampler(cfg[\"candidate_sampling\"])\nsubmission_builder = SubmissionBuilder()\n\nprint(\"Pipeline modules initialised ✓\")\n","outputs":[],"execution_count":null},{"cell_type":"code","metadata":{},"source":"# ────────────────────────────────────────────────────────────\n# Main Prediction Loop\n# ────────────────────────────────────────────────────────────\nimport time\nimport logging\nlogging.basicConfig(level=logging.INFO,\n                    format=\"%(asctime)s [%(levelname)s] %(message)s\")\nlogger = logging.getLogger(\"notebook\")\n\nall_predictions = []\nt_total = time.time()\n\nfor idx, row in test_sequences.iterrows():\n    target_id     = row[\"target_id\"]\n    sequence      = row[\"sequence\"]\n    seq_len       = int(row[\"seq_len\"])\n    stoichiometry = str(row.get(\"stoichiometry\", \"A:1\"))\n    has_ligands   = bool(row.get(\"has_ligands\", False))\n\n    logger.info(f\"[{idx+1}/{len(test_sequences)}] {target_id}  \"\n                f\"len={seq_len}  stoich={stoichiometry}\")\n\n    # Validate sequence\n    if not sequence or not all(c in \"ACGUN\" for c in sequence.upper()):\n        logger.warning(f\"  Skipping {target_id}: invalid characters\")\n        continue\n\n    try:\n        t_seq = time.time()\n\n        # Stage 1: Secondary structure\n        sec_struct = ss_predictor.predict(sequence)\n\n        # Stage 2: Family classification\n        family = family_clf.classify(sequence, sec_struct)\n        logger.info(f\"  Family: {family.name}\")\n\n        # Stage 3: Template search\n        templates = template_searcher.search(sequence, family)\n        logger.info(f\"  Templates: {len(templates)}, \"\n                    f\"best TM≈{templates[0].expected_tm:.3f}\" if templates else \"  Templates: 0\")\n\n        # Stage 4: Route\n        branch = router.route(templates, family)\n        logger.info(f\"  Branch: {branch}\")\n\n        # Stage 5+6+7: Predict → correct → rank\n        raw_structs = candidate_sampler.sample(\n            sequence=sequence, sec_struct=sec_struct,\n            templates=templates if branch == \"tbm\" else [],\n            predictor=structure_pred, branch=branch,\n            target_id=target_id,\n        )\n        corrected = [motif_corrector.correct(s, sec_struct) for s in raw_structs]\n        ranked    = candidate_sampler.rank(corrected)\n\n        logger.info(f\"  Done in {time.time()-t_seq:.1f}s  \"\n                    f\"pLDDT={ranked[0].plddt:.1f}\")\n\n        all_predictions.append({\n            \"target_id\":     target_id,\n            \"sequence\":      sequence,\n            \"stoichiometry\": stoichiometry,\n            \"has_ligands\":   has_ligands,\n            \"family\":        family.name,\n            \"branch\":        branch,\n            \"n_templates\":   len(templates),\n            \"structures\":    ranked,\n        })\n\n    except Exception as e:\n        logger.error(f\"  FAILED {target_id}: {e}\", exc_info=True)\n        all_predictions.append({\n            \"target_id\":  target_id,\n            \"sequence\":   sequence,\n            \"stoichiometry\": stoichiometry,\n            \"has_ligands\":  False,\n            \"family\":     \"error\",\n            \"branch\":     \"fallback\",\n            \"n_templates\": 0,\n            \"structures\": candidate_sampler.make_fallback(sequence),\n        })\n\nlogger.info(f\"\\nAll {len(all_predictions)} sequences done in \"\n            f\"{(time.time()-t_total)/60:.1f} min\")\n","outputs":[],"execution_count":null},{"cell_type":"code","metadata":{},"source":"# ────────────────────────────────────────────────────────────\n# Build submission.csv\n# ────────────────────────────────────────────────────────────\noutput_path = f\"{OUTPUT_DIR}/submission.csv\"\nsubmission_builder.build(all_predictions, output_path)\n\n# Validate format\nimport pandas as pd\ndf = pd.read_csv(output_path)\nprint(f\"\\nsubmission.csv written: {output_path}\")\nprint(f\"  Rows    : {len(df):,}\")\nprint(f\"  Columns : {list(df.columns)}\")\nprint(f\"  Targets : {df['ID'].str.rsplit('_',n=1).str[0].nunique()}\")\nprint()\nprint(df.head(3).to_string())\n","outputs":[],"execution_count":null},{"cell_type":"code","metadata":{},"source":"# ────────────────────────────────────────────────────────────\n# Local Validation (skipped on Kaggle)\n# ────────────────────────────────────────────────────────────\n# Run only locally to score against validation_labels.csv\nif not KAGGLE_ENV:\n    import subprocess, sys\n    labels_path = f\"{DATA_DIR}/validation_labels.csv\"\n    if Path(labels_path).exists():\n        result = subprocess.run(\n            [sys.executable, \"scripts/validate_submission.py\",\n             \"--submission\", output_path,\n             \"--labels\",     labels_path],\n            capture_output=False\n        )\n    else:\n        print(f\"Labels not found at {labels_path} — skipping validation\")\nelse:\n    print(\"On Kaggle: validation scoring not available (no labels)\")\n","outputs":[],"execution_count":null}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.0"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","datasetId":"stanford-rna-3d-folding-2"}],"isGpuEnabled":true,"isInternetEnabled":false}},"nbformat":4,"nbformat_minor":5}