{"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":15231210},{"sourceType":"datasetVersion","sourceId":14874339,"datasetId":9502242,"databundleVersionId":15736806},{"sourceType":"datasetVersion","sourceId":10855324,"datasetId":6742586,"databundleVersionId":11219268},{"sourceType":"kernelVersion","sourceId":299291353}],"dockerImageVersionId":31287,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":265.887498,"end_time":"2026-03-12T21:18:19.583777","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2026-03-12T21:13:53.696279","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport sys\nimport json\nimport time\nimport gc\nimport re\nimport warnings\nfrom pathlib import Path\nimport pandas as pd\nimport numpy as np\nimport torch\nfrom Bio.Align import PairwiseAligner\nfrom numba import njit\nfrom tqdm import tqdm\nimport subprocess\nimport pickle\n\ntry:\n    from rdkit import Chem\n    from rdkit.Chem import AllChem\n    from rdkit.Chem import rdFingerprintGenerator\n    from rdkit import DataStructs\n    RDKIT_AVAILABLE = True\n    mfpgen = rdFingerprintGenerator.GetMorganGenerator(radius=2, fpSize=1024)\nexcept ImportError:\n    RDKIT_AVAILABLE = False\n    print(\"Warning: RDKit not found. Ligand-guided selection disabled.\")\n\n# ViennaRNA for thermodynamic ranking\ntry:\n    import RNA\n    VIENNA_AVAILABLE = True\nexcept ImportError:\n    VIENNA_AVAILABLE = False\n    print(\"Warning: ViennaRNA not found. Thermodynamic ranking will be skipped.\")\n\n# Determinism hygiene\nos.environ[\"PYTHONHASHSEED\"] = \"0\"\nos.environ[\"LAYERNORM_TYPE\"] = \"torch\"\nos.environ.setdefault(\"RNA_MSA_DEPTH_LIMIT\", \"512\")\nwarnings.filterwarnings(\"ignore\")\n\n# ─────────────── Paths & Constants ───────────────────────────────────────────\nDATA_BASE              = \"/kaggle/input/competitions/stanford-rna-3d-folding-2\"\nMSA_DIR                = f\"{DATA_BASE}/MSA\"\nMETADATA_CSV           = f\"{DATA_BASE}/extra/rna_metadata.csv\"\nPDB_SEQRES_FASTA       = f\"{DATA_BASE}/PDB_RNA/pdb_seqres_NA.fasta\"\nPDB_RNA_CIF_DIR        = f\"{DATA_BASE}/PDB_RNA\"                        # NEW: CIF 3D structures\nPDB_RELEASE_DATES_CSV  = f\"{DATA_BASE}/PDB_RNA/pdb_release_dates_NA.csv\"  # NEW: temporal filter\nDEFAULT_TEST_CSV       = f\"{DATA_BASE}/test_sequences.csv\"\nDEFAULT_TRAIN_CSV      = f\"{DATA_BASE}/train_sequences.csv\"\nDEFAULT_TRAIN_LBLS     = f\"{DATA_BASE}/train_labels.csv\"\nDEFAULT_VAL_CSV        = f\"{DATA_BASE}/validation_sequences.csv\"\nDEFAULT_VAL_LBLS       = f\"{DATA_BASE}/validation_labels.csv\"\nDEFAULT_OUTPUT         = \"/kaggle/working/submission.csv\"\n\n# Partner molecule size cap (residues) — prevents OOM on T4 16GB\nMAX_PARTNER_RESIDUES   = 500\nMAX_TOTAL_TOKENS       = 800\n\nDEFAULT_CODE_DIR = (\n    \"/kaggle/input/datasets/qiweiyin/protenix-v1-adjusted\"\n    \"/Protenix-v1-adjust-v2/Protenix-v1-adjust-v2/Protenix-v1\"\n)\nDEFAULT_ROOT_DIR = DEFAULT_CODE_DIR\n\nMODEL_NAME    = \"protenix_base_20250630_v1.0.0\"\nN_SAMPLE      = 5\nSEED          = 108\n\n# ── Inference precision & compute knobs ─────────────────────────────────────\n# dtype: \"bf16\" halves VRAM vs \"fp32\" with negligible quality loss on T4/A100.\nDTYPE         = os.environ.get(\"INFER_DTYPE\",        \"bf16\")\n# N_CYCLE: Pairformer recycling passes. 10 = full quality; 4 = ~40% faster.\nN_CYCLE       = int(os.environ.get(\"INFER_N_CYCLE\",  \"12\"))\n# N_STEP: Diffusion denoising steps. 200 = best quality; 20 = fast draft.\nN_STEP        = int(os.environ.get(\"INFER_N_STEP\",   \"100\"))\n# Triangle kernels: \"cuequivariance\" saves ~20% VRAM vs \"torch\" on CUDA >= 11.8\n# Actual CLI flags: --triangle_multiplicative and --triangle_attention\nTRIMUL_KERNEL = os.environ.get(\"INFER_TRIMUL\",       \"cuequivariance\")\nTRIATT_KERNEL = os.environ.get(\"INFER_TRIATT\",       \"cuequivariance\")\n\n# SLIDING WINDOW CONSTANTS\nMAX_SEQ_LEN   = 512\nCHUNK_OVERLAP = 64\n\nMIN_SIMILARITY       = float(os.environ.get(\"MIN_SIMILARITY\",       \"0.0\"))\nMIN_PERCENT_IDENTITY = float(os.environ.get(\"MIN_PERCENT_IDENTITY\", \"50.0\"))\n\n# Innovation: Re-threading threshold (sequences below this %ID get loop-level fragment assembly)\nRETHREAD_IDENTITY_THRESHOLD = 80.0\n\nUSE_PROTENIX  = True\nUSE_CODOCK    = True   # CoDock ligand-pocket scoring for sample re-ranking\n\ndef parse_bool(value: str, default: bool = False) -> str:\n    v = str(value).strip().lower()\n    if v in {\"1\", \"true\", \"t\", \"yes\", \"y\", \"on\"}: return \"true\"\n    if v in {\"0\", \"false\", \"f\", \"no\", \"n\", \"off\"}: return \"false\"\n    return \"true\" if default else \"false\"\n\nUSE_MSA      = parse_bool(os.environ.get(\"USE_MSA\",      \"true\"))   # NEW: was false — enables Protenix MSA co-evolution\nUSE_TEMPLATE = parse_bool(os.environ.get(\"USE_TEMPLATE\", \"true\"))   # CIF files available via PDB_RNA_CIF_DIR\nUSE_RNA_MSA  = parse_bool(os.environ.get(\"USE_RNA_MSA\",  \"true\"))\nMODEL_N_SAMPLE = int(os.environ.get(\"MODEL_N_SAMPLE\", str(N_SAMPLE)))\n\nIS_KAGGLE = bool(os.environ.get(\"KAGGLE_IS_COMPETITION_RERUN\", \"\"))\nLOCAL_N_SAMPLES = 2\n\n# ─── CCD ion codes: these must use the \"ion\" JSON key, never \"ligand\" + SMILES ─\n# The json_parser's build_ligand() routes the \"ion\" key directly to CCD lookup\n# (no RDKit conformer generation). Routing these through {\"ligand\": SMILES} causes\n# EmbedMolecule to be called on single atoms / metal complexes → guaranteed crash.\nKNOWN_ION_CCD_CODES: frozenset = frozenset({\n    # Alkali metals\n    \"LI\", \"NA\", \"K\", \"RB\", \"CS\",\n    # Alkaline earth\n    \"MG\", \"CA\", \"SR\", \"BA\",\n    # Transition / heavy metals common in RNA structures\n    \"MN\", \"FE\", \"FE2\", \"CO\", \"NI\", \"CU\", \"CU1\", \"ZN\", \"CD\", \"HG\",\n    \"AU\", \"AG\", \"PB\",\n    # Halide ions\n    \"F\", \"CL\", \"BR\", \"IOD\",\n    # Polyatomic inorganic ions\n    \"SO4\", \"PO4\", \"NO3\", \"CO3\",\n    # Hexaammine metal complexes (IRI = hexaammineiridium, COB = hexaamminecobalts)\n    # These have SMILES like [NH3]->[Ir+3] that crash EmbedMolecule.\n    # They must go through the CCD path as ions.\n    \"IRI\", \"COB\",\n})\n\nif IS_KAGGLE:\n    print(\"Running in KAGGLE COMPETITION mode — all test targets will be processed.\")\nelse:\n    print(f\"Running in LOCAL mode — only {LOCAL_N_SAMPLES} targets for quick testing.\")\n\ndef seed_everything(seed: int) -> None:\n    os.environ[\"CUBLAS_WORKSPACE_CONFIG\"] = \":4096:8\"\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    np.random.seed(seed)\n    torch.backends.cudnn.benchmark = False\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.enabled = True\n    torch.use_deterministic_algorithms(True)\n\ndef resolve_paths():\n    test_csv   = os.environ.get(\"TEST_CSV\",           DEFAULT_TEST_CSV)\n    output_csv = os.environ.get(\"SUBMISSION_CSV\",     DEFAULT_OUTPUT)\n    code_dir   = os.environ.get(\"PROTENIX_CODE_DIR\",  DEFAULT_CODE_DIR)\n    root_dir   = os.environ.get(\"PROTENIX_ROOT_DIR\",  DEFAULT_ROOT_DIR)\n    return test_csv, output_csv, code_dir, root_dir\n\ndef ensure_required_files(root_dir: str) -> None:\n    for p, name in [\n        (Path(root_dir) / \"checkpoint\" / f\"{MODEL_NAME}.pt\",          \"checkpoint\"),\n        (Path(root_dir) / \"common\" / \"components.cif\",                \"CCD file\"),\n        (Path(root_dir) / \"common\" / \"components.cif.rdkit_mol.pkl\",  \"CCD cache\"),\n    ]:\n        if not p.exists(): raise FileNotFoundError(f\"Missing {name}: {p}\")\n\ndef ensure_mmcif_dir(cif_source: str) -> str:\n    \"\"\"\n    Protenix needs a writable mmcif directory.\n    /kaggle/input is read-only, so we create /kaggle/working/mmcif and\n    populate it with symlinks to every .cif / .cif.gz file in cif_source.\n    Returns the writable mmcif dir path to pass into Protenix configs.\n    \"\"\"\n    writable_mmcif = \"/kaggle/working/mmcif\"\n\n    if os.path.isdir(writable_mmcif):\n        n = len(os.listdir(writable_mmcif))\n        print(f\"✓ mmcif dir already present at {writable_mmcif} ({n} files)\")\n        return writable_mmcif\n\n    os.makedirs(writable_mmcif, exist_ok=True)\n\n    cif_files = (\n        list(Path(cif_source).glob(\"*.cif\")) +\n        list(Path(cif_source).glob(\"*.cif.gz\"))\n    )\n    linked = 0\n    for cif_file in cif_files:\n        link_path = Path(writable_mmcif) / cif_file.name\n        if not link_path.exists():\n            os.symlink(str(cif_file), str(link_path))\n            linked += 1\n\n    print(f\"✓ mmcif dir ready at {writable_mmcif} ({linked} CIF files symlinked from {cif_source})\")\n    return writable_mmcif\n\n# ─────────────── Ligand Engine & 2D Folding ───────────────────────────────\ndef extract_ligand_features(smiles_str, ligand_ids_str=\"\"):\n    \"\"\"\n    Compute RDKit fingerprints only for genuine organic ligands.\n    Ions (identified by CCD code in ligand_ids_str) are skipped — they carry\n    no meaningful Tanimoto-comparable fingerprint.\n    \"\"\"\n    if not RDKIT_AVAILABLE or pd.isna(smiles_str) or str(smiles_str).strip() == \"\": return [], 0\n    fps, max_heavy_atoms = [], 0\n\n    smi_list = [s.strip() for s in str(smiles_str).split(\";\")]\n    ccd_list = [c.strip().upper() for c in str(ligand_ids_str).split(\";\")] if ligand_ids_str else [\"\"] * len(smi_list)\n    # Pad ccd_list if lengths differ\n    while len(ccd_list) < len(smi_list):\n        ccd_list.append(\"\")\n\n    for ccd_code, clean_smi in zip(ccd_list, smi_list):\n        # Skip ions — they contribute no useful structural fingerprint\n        if ccd_code in KNOWN_ION_CCD_CODES:\n            continue\n        if not clean_smi:\n            continue\n        mol = Chem.MolFromSmiles(clean_smi)\n        if mol:\n            heavy_atoms = mol.GetNumHeavyAtoms()\n            if heavy_atoms > max_heavy_atoms: max_heavy_atoms = heavy_atoms\n            if heavy_atoms < 3: continue\n            fp = mfpgen.GetFingerprint(mol)\n            fps.append(fp)\n    return fps, max_heavy_atoms\n\ndef calc_max_ligand_similarity(query_fps, template_fps):\n    if not query_fps or not template_fps: return 0.0\n    max_sim = 0.0\n    for q_fp in query_fps:\n        for t_fp in template_fps:\n            sim = DataStructs.TanimotoSimilarity(q_fp, t_fp)\n            if sim > max_sim: max_sim = sim\n    return max_sim\n\ndef get_rna_base_pairs(seq: str) -> list:\n    try:\n        import RNA\n        fc = RNA.fold_compound(seq)\n        ss, _ = fc.mfe()\n        stack, pairs = [], []\n        for i, char in enumerate(ss):\n            if char == '(': stack.append(i)\n            elif char == ')':\n                if stack: pairs.append((stack.pop(), i))\n        return pairs\n    except ImportError:\n        return []\n\n# ─────────────── Innovation 4: ΔG Thermodynamic Ranking ──────────────────────\ndef compute_delta_g(seq: str) -> float:\n    \"\"\"\n    Compute the MFE free energy (ΔG) for a sequence using ViennaRNA.\n    Lower values = more thermodynamically stable = physically more feasible.\n    Returns 0.0 if ViennaRNA is not available.\n    \"\"\"\n    if not VIENNA_AVAILABLE or not seq:\n        return 0.0\n    try:\n        fc = RNA.fold_compound(seq)\n        _, mfe = fc.mfe()\n        return float(mfe)\n    except Exception:\n        return 0.0\n\ndef thermodynamic_rank_samples(samples: list, seq: str, plddt_scores: list = None) -> list:\n    \"\"\"\n    Innovation: Rank Protenix samples using a combined ΔG + pLDDT score.\n    ΔG measures physical feasibility; pLDDT measures model confidence.\n    We compute ΔG for the sequence once (it's sequence-level) and use it as\n    a tiebreaker to prefer samples that are structurally compact.\n    \n    For structure-level ranking we use radius of gyration (Rg) as a structural\n    proxy for thermodynamic compactness, biased by the Debye-Hückel GC effect.\n    \"\"\"\n    if not samples:\n        return samples\n\n    gc_content = sum(1 for b in seq if b in ('G', 'C')) / max(len(seq), 1)\n    delta_g = compute_delta_g(seq)\n\n    def score_sample(idx):\n        coords = samples[idx]\n        centroid = coords.mean(axis=0)\n        rg = np.sqrt(np.mean(np.sum((coords - centroid)**2, axis=1)))\n        # High GC content → tighter Rg is more physical (Debye-Hückel screening)\n        rg_target = 10.0 * (len(seq) ** 0.38) * (1.0 - 0.15 * gc_content)\n        rg_penalty = abs(rg - rg_target)\n        plddt = plddt_scores[idx] if plddt_scores and idx < len(plddt_scores) else 0.5\n        # Combined: minimize rg_penalty, maximize pLDDT, weight by |ΔG| (more negative = better)\n        combined = -plddt + 0.01 * rg_penalty - 0.001 * abs(delta_g)\n        return combined\n\n    ranked_indices = sorted(range(len(samples)), key=score_sample)\n    return [samples[i] for i in ranked_indices]\n\n# ─────────────── CoDock: Ligand-Pocket Proximity Scoring ─────────────────────\ndef codock_score_sample(rna_c1_coords: np.ndarray, smiles_list: list) -> float:\n    \"\"\"\n    CoDock Finalization: given an RNA backbone (C1' atoms only) and a list of\n    ligand SMILES, estimate how well the RNA pocket accommodates the ligands.\n\n    Strategy (no extra dependencies — uses RDKit already in scope):\n      1. For each organic SMILES, generate N_CONF ETKDG conformers.\n      2. For each conformer, place its centroid at the RNA pocket centre and\n         compute a contact score: sum of 1/(d+0.5)^2 for all ligand-heavy ×\n         C1' pairs within CONTACT_CUTOFF Å, minus a clash penalty for d < CLASH_THR Å.\n      3. Keep the best-scoring conformer pose; sum scores across all ligands.\n\n    A higher score = better-shaped pocket for the ligand = more physically\n    plausible RNA structure.  Used as a third ranking signal alongside pLDDT\n    and ΔG thermodynamic ranking.\n\n    Returns a float (higher = better).  Returns 0.0 if RDKit unavailable or\n    no valid organic SMILES are present.\n    \"\"\"\n    if not RDKIT_AVAILABLE or not smiles_list:\n        return 0.0\n\n    from rdkit.Chem import AllChem\n\n    N_CONF         = 10       # conformers per ligand\n    CONTACT_CUTOFF = 15.0     # Å — RNA atoms beyond this are ignored\n    CLASH_THR      = 2.5      # Å — heavy penalty if any atom closer than this\n    CLASH_PENALTY  = 5.0      # subtracted per clashing pair\n\n    # RNA pocket centre = centroid of all C1' atoms\n    pocket_centre = rna_c1_coords.mean(axis=0)\n    rna_arr = rna_c1_coords.astype(np.float32)\n\n    total_score = 0.0\n\n    for smi in smiles_list:\n        # Guard: SMILES may be NaN float if column was empty in the CSV\n        if not isinstance(smi, str):\n            continue\n        smi = smi.strip()\n        if not smi or smi.lower() == \"nan\":\n            continue\n        mol = Chem.MolFromSmiles(smi)\n        if mol is None or mol.GetNumHeavyAtoms() < 3:\n            continue\n        mol_h = Chem.AddHs(mol)\n\n        # Generate conformers\n        params = AllChem.ETKDGv3()\n        params.randomSeed = 42\n        params.numThreads = 1\n        ids = AllChem.EmbedMultipleConfs(mol_h, numConfs=N_CONF, params=params)\n        if len(ids) == 0:\n            # Fallback: random coords\n            if AllChem.EmbedMolecule(mol_h, useRandomCoords=True) != 0:\n                continue\n            ids = [0]\n\n        best_conf_score = -1e9\n        for cid in ids:\n            conf = mol_h.GetConformer(cid)\n            lig_pos = np.array([list(conf.GetAtomPosition(a.GetIdx()))\n                                for a in mol_h.GetAtoms()\n                                if a.GetAtomicNum() > 1], dtype=np.float32)\n            if len(lig_pos) == 0:\n                continue\n\n            # Translate ligand centroid → RNA pocket centre\n            lig_pos += (pocket_centre - lig_pos.mean(axis=0))\n\n            # Vectorised distance matrix: (n_lig, n_rna)\n            diff = lig_pos[:, None, :] - rna_arr[None, :, :]   # (L, R, 3)\n            dists = np.sqrt((diff ** 2).sum(axis=-1))            # (L, R)\n\n            # Contact score: attractive term\n            mask_contact = dists < CONTACT_CUTOFF\n            score = np.sum(1.0 / (dists[mask_contact] + 0.5) ** 2)\n\n            # Clash penalty: repulsive term\n            n_clashes = np.sum(dists < CLASH_THR)\n            score -= n_clashes * CLASH_PENALTY\n\n            if score > best_conf_score:\n                best_conf_score = score\n\n        total_score += max(best_conf_score, 0.0)\n\n    return float(total_score)\n\n\ndef codock_rank_samples(\n    samples:      list,\n    seq:          str,\n    smiles_list:  list,\n    plddt_scores: list = None,\n    delta_g:      float = 0.0,\n) -> list:\n    \"\"\"\n    Re-rank Protenix samples using a combined three-signal score:\n      1. pLDDT  — model confidence (higher = better)\n      2. ΔG     — thermodynamic stability (more negative = better)\n      3. CoDock — ligand pocket-fit score (higher = better pocket geometry)\n\n    Returns samples sorted best-first.\n    \"\"\"\n    if not samples:\n        return samples\n\n    def combined_score(idx):\n        plddt   = plddt_scores[idx] if plddt_scores and idx < len(plddt_scores) else 0.5\n        dock    = codock_score_sample(samples[idx], smiles_list) if USE_CODOCK else 0.0\n        # Normalise each signal: negative = better for minimisation sort\n        return -plddt - 0.002 * dock - 0.001 * abs(delta_g)\n\n    ranked = sorted(range(len(samples)), key=combined_score)\n    return [samples[i] for i in ranked]\n\n\n# ─────────────── KABSCH 3D ASSEMBLY ALGORITHM ─────────────────────────────\ndef kabsch_alignment(P, Q):\n    \"\"\"Calculates optimal rotation and translation to map Matrix P onto Matrix Q using SVD.\"\"\"\n    P_centroid = np.mean(P, axis=0)\n    Q_centroid = np.mean(Q, axis=0)\n    P_centered = P - P_centroid\n    Q_centered = Q - Q_centroid\n\n    H = P_centered.T @ Q_centered\n    U, S, Vt = np.linalg.svd(H)\n    R = Vt.T @ U.T\n\n    # Correct for improper rotations (reflections)\n    if np.linalg.det(R) < 0:\n        Vt[2, :] *= -1\n        R = Vt.T @ U.T\n\n    t = Q_centroid - P_centroid @ R.T\n    return R, t\n\ndef stitch_chunks(chunks, overlap=64):\n    \"\"\"Assembles sliding-window chunks via 3D point-cloud registration and alpha blending.\"\"\"\n    if not chunks: return np.zeros((0, 3))\n    if len(chunks) == 1: return chunks[0]\n\n    assembled = chunks[0].copy()\n    \n    for i in range(1, len(chunks)):\n        next_chunk = chunks[i].copy()\n        \n        # P = Next chunk's anchor. Q = Assembled molecule's tail.\n        P = next_chunk[:overlap]\n        Q = assembled[-overlap:]\n        \n        # 1. Kabsch Spatial Registration\n        R, t = kabsch_alignment(P, Q)\n        aligned_next = (next_chunk @ R.T) + t\n        \n        # 2. Linear Interpolation (Alpha Blending) of the Overlap Region\n        weights_next = np.linspace(0, 1, overlap)[:, None]\n        weights_prev = 1.0 - weights_next\n        blended_overlap = assembled[-overlap:] * weights_prev + aligned_next[:overlap] * weights_next\n        \n        # 3. Stitching\n        assembled[-overlap:] = blended_overlap\n        assembled = np.vstack([assembled, aligned_next[overlap:]])\n        \n    return assembled\n\ndef sort_by_centroid(coords: np.ndarray) -> np.ndarray:\n    n_samples = coords.shape[0]\n    if n_samples <= 1: return coords\n    rmsds = np.zeros((n_samples, n_samples))\n    for i in range(n_samples):\n        for j in range(i + 1, n_samples):\n            diff = coords[i] - coords[j]\n            rmsd = np.sqrt(np.mean(np.sum(diff**2, axis=-1)))\n            rmsds[i, j] = rmsds[j, i] = rmsd\n    avg_rmsd = rmsds.sum(axis=1) / (n_samples - 1)\n    sorted_indices = np.argsort(avg_rmsd)\n    return coords[sorted_indices]\n\n\n# ═══════════════════════════════════════════════════════════════════════════\n# Vfold-Inspired Hybrid Pipeline\n# ─────────────────────────────────────────────────────────────────────────\n# 1. VfoldScaffoldGenerator  — ViennaRNA MFE -> A-form helices + interpolated\n#    loops -> coarse-grained global 3D scaffold for the full sequence.\n# 2. write_chunk_vfold_cif   — Crops the scaffold to a chunk window and\n#    writes a Protenix-compatible mmCIF into mmcif_dir. Protenix template\n#    featurizer discovers it via 100% sequence-identity match (top template).\n# 3. align_assembled_to_vfold — After Protenix assembles chunks, performs a\n#    global Kabsch alignment onto the Vfold scaffold to eliminate drift\n#    accumulated by sequential stitch_chunks calls on long sequences.\n# ═══════════════════════════════════════════════════════════════════════════\n\n_CIF_HEADER_FIELDS = [\n    \"loop_\",\n    \"_atom_site.group_PDB\",\n    \"_atom_site.id\",\n    \"_atom_site.type_symbol\",\n    \"_atom_site.label_atom_id\",\n    \"_atom_site.label_alt_id\",\n    \"_atom_site.label_comp_id\",\n    \"_atom_site.label_asym_id\",\n    \"_atom_site.label_entity_id\",\n    \"_atom_site.label_seq_id\",\n    \"_atom_site.pdbx_PDB_ins_code\",\n    \"_atom_site.Cartn_x\",\n    \"_atom_site.Cartn_y\",\n    \"_atom_site.Cartn_z\",\n    \"_atom_site.occupancy\",\n    \"_atom_site.B_iso_or_equiv\",\n    \"_atom_site.auth_seq_id\",\n    \"_atom_site.auth_comp_id\",\n    \"_atom_site.auth_asym_id\",\n    \"_atom_site.pdbx_PDB_model_num\",\n]\n\n\ndef _parse_dot_bracket(dot_bracket: str) -> dict:\n    \"\"\"Return {i: j} Watson-Crick pair map (0-indexed).\"\"\"\n    stack, pairs = [], {}\n    for i, ch in enumerate(dot_bracket):\n        if ch == \"(\":\n            stack.append(i)\n        elif ch == \")\" and stack:\n            j = stack.pop()\n            pairs[j] = i\n            pairs[i] = j\n    return pairs\n\n\ndef _group_into_helix_runs(pair_map: dict) -> list:\n    \"\"\"\n    Extract maximal consecutive stems: runs of pairs\n    (i,j),(i+1,j-1),(i+2,j-2),...\n    Returns list of [(i,j),...] runs (minimum 2 pairs per run).\n    \"\"\"\n    visited, runs = set(), []\n    for i in sorted(pair_map):\n        j = pair_map[i]\n        if i >= j or i in visited:\n            continue\n        run = []\n        ii, jj = i, j\n        while (ii in pair_map and pair_map[ii] == jj and ii not in visited):\n            run.append((ii, jj))\n            visited.add(ii)\n            visited.add(jj)\n            ii += 1\n            jj -= 1\n        if len(run) >= 2:\n            runs.append(run)\n    return runs\n\n\ndef _build_aform_coords(run: list) -> dict:\n    \"\"\"\n    Compute A-form helix C1-prime positions for one stem run.\n    A-form parameters: rise=2.81 A, twist=32.7 deg, radius=8.7 A.\n    Returns {residue_idx: np.array([x,y,z])}.\n    \"\"\"\n    RISE   = 2.81\n    TWIST  = np.radians(32.7)\n    RADIUS = 8.7\n    out = {}\n    for step, (i, j) in enumerate(run):\n        z    = step * RISE\n        ang1 = step * TWIST\n        ang2 = np.pi + step * TWIST\n        out[i] = np.array([RADIUS * np.cos(ang1), RADIUS * np.sin(ang1), z])\n        out[j] = np.array([RADIUS * np.cos(ang2), RADIUS * np.sin(ang2), z])\n    return out\n\n\ndef _fill_chain_coords(partial: dict, seq_len: int) -> np.ndarray:\n    \"\"\"\n    Fill every position via linear interpolation between placed anchors;\n    extrapolate ends along Z with 2.81 A rise per residue.\n    Returns (seq_len, 3) float64 array.\n    \"\"\"\n    RISE = 2.81\n    coords = np.zeros((seq_len, 3), dtype=np.float64)\n    placed = np.zeros(seq_len, dtype=bool)\n    for idx, xyz in partial.items():\n        if 0 <= idx < seq_len:\n            coords[idx] = xyz\n            placed[idx] = True\n    if not placed.any():\n        for k in range(seq_len):\n            ang = k * np.radians(32.7)\n            coords[k] = [8.7 * np.cos(ang), 8.7 * np.sin(ang), k * RISE]\n        return coords\n    pidx = np.where(placed)[0]\n    for k in range(pidx[0]):\n        coords[k] = coords[pidx[0]] - np.array([0, 0, (pidx[0] - k) * RISE])\n    placed[:pidx[0]] = True\n    for k in range(pidx[-1] + 1, seq_len):\n        coords[k] = coords[pidx[-1]] + np.array([0, 0, (k - pidx[-1]) * RISE])\n    placed[pidx[-1] + 1:] = True\n    i = 0\n    while i < seq_len:\n        if not placed[i]:\n            right = i\n            while right < seq_len and not placed[right]:\n                right += 1\n            left = i - 1\n            gap  = right - left\n            for k in range(left + 1, right):\n                t = (k - left) / gap\n                coords[k] = coords[left] * (1 - t) + coords[right] * t\n            placed[i:right] = True\n            i = right\n        else:\n            i += 1\n    return coords\n\n\ndef _write_vfold_cif(seq, coords, chain_segs, pdb_id, out_path):\n    \"\"\"Write a Protenix-compatible mmCIF with C1-prime atoms only.\"\"\"\n    lines = [\n        \"data_\" + pdb_id, \"#\",\n        \"_entry.id \" + pdb_id, \"#\",\n    ] + _CIF_HEADER_FIELDS\n    res_info = {}\n    for eid, (cid, cs, ce) in enumerate(chain_segs, 1):\n        for li in range(ce - cs):\n            res_info[cs + li] = (cid, eid, li + 1)\n    atom_id = 1\n    for gi in range(min(len(seq), len(coords))):\n        cid, eid, lseq = res_info.get(gi, (\"A\", 1, gi + 1))\n        rc = seq[gi].upper()\n        x = float(coords[gi, 0])\n        y = float(coords[gi, 1])\n        z = float(coords[gi, 2])\n        # C1-prime contains apostrophe — double-quote in CIF\n        lines.append(\n            'ATOM ' + str(atom_id) + ' C \"C1\\'\" . ' +\n            rc + ' ' + cid + ' ' + str(eid) + ' ' + str(lseq) + ' ? ' +\n            '{:.3f} {:.3f} {:.3f}'.format(x, y, z) + ' 1.00 0.00 ' +\n            str(lseq) + ' ' + rc + ' ' + cid + ' 1'\n        )\n        atom_id += 1\n    lines.append(\"#\")\n    with open(out_path, \"w\") as fh:\n        fh.write(\"\\n\".join(lines) + \"\\n\")\n\n\nclass VfoldScaffoldGenerator:\n    \"\"\"\n    Global coarse-grained C1-prime 3D scaffold for a full multi-chain RNA.\n\n    Stems  -> standard A-form helices (rise 2.81 A, twist 32.7 deg).\n    Loops  -> linearly interpolated to maintain backbone connectivity.\n    Chains -> separated by CHAIN_SEP A along X to avoid clashes.\n\n    Usage:\n        gen  = VfoldScaffoldGenerator(sequence, strands)\n        gen.generate()\n        path = gen.write_cif(pdb_id, mmcif_dir)\n        arr  = gen.chunk_coords(start, end)   # (chunk_len, 3)\n    \"\"\"\n    CHAIN_SEP      = 60.0\n    _CHAIN_LETTERS = \"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz\"\n\n    def __init__(self, sequence: str, strands: list):\n        self.sequence = sequence\n        self.strands  = strands\n        segs, offset, ci = [], 0, 0\n        for chain_seq, count in strands:\n            for _ in range(count):\n                cid = self._CHAIN_LETTERS[ci % len(self._CHAIN_LETTERS)]\n                segs.append((cid, offset, offset + len(chain_seq)))\n                offset += len(chain_seq)\n                ci     += 1\n        self.chain_segs   = segs\n        self._coord_array = np.zeros((len(sequence), 3), dtype=np.float64)\n        self._generated   = False\n\n    def generate(self):\n        \"\"\"Run ViennaRNA MFE + A-form placement for all chains.\"\"\"\n        x_off = 0.0\n        for cid, cstart, cend in self.chain_segs:\n            chain_seq = self.sequence[cstart:cend]\n            chain_len = len(chain_seq)\n            if VIENNA_AVAILABLE:\n                fc    = RNA.fold_compound(chain_seq)\n                db, _ = fc.mfe()\n            else:\n                db = \".\" * chain_len\n            pair_map = _parse_dot_bracket(db)\n            runs     = _group_into_helix_runs(pair_map)\n            partial  = {}\n            for run in runs:\n                partial.update(_build_aform_coords(run))\n            chain_c = _fill_chain_coords(partial, chain_len)\n            chain_c[:, 0] += x_off\n            x_off         += self.CHAIN_SEP\n            self._coord_array[cstart:cend] = chain_c\n        self._generated = True\n        return self\n\n    def write_cif(self, pdb_id: str, out_dir: str) -> str:\n        \"\"\"Write full-sequence scaffold CIF. Returns path.\"\"\"\n        if not self._generated:\n            self.generate()\n        os.makedirs(out_dir, exist_ok=True)\n        path = os.path.join(out_dir, pdb_id + \".cif\")\n        _write_vfold_cif(self.sequence, self._coord_array,\n                         self.chain_segs, pdb_id, path)\n        return path\n\n    def chunk_coords(self, chunk_start: int, chunk_end: int) -> np.ndarray:\n        \"\"\"Return (chunk_len, 3) coordinate slice.\"\"\"\n        if not self._generated:\n            self.generate()\n        return self._coord_array[chunk_start:chunk_end].copy()\n\n\ndef write_chunk_vfold_cif(gen, chunk_id, chunk_start, chunk_end, mmcif_dir):\n    \"\"\"\n    Crop the global Vfold scaffold to [chunk_start, chunk_end) and write a\n    Protenix-compatible mmCIF into mmcif_dir.\n\n    Protenix template featurizer discovers this via sequence-identity search;\n    the CIF carries the exact chunk sequence (100% identity = top template).\n\n    Returns (cif_path, pdb_id).\n    \"\"\"\n    import hashlib\n    if not gen._generated:\n        gen.generate()\n    chunk_seq    = gen.sequence[chunk_start:chunk_end]\n    chunk_coords = gen.chunk_coords(chunk_start, chunk_end)\n    chunk_chain_segs = []\n    for cid, cstart, cend in gen.chain_segs:\n        ov_s = max(chunk_start, cstart) - chunk_start\n        ov_e = min(chunk_end,   cend)   - chunk_start\n        if ov_s < ov_e:\n            chunk_chain_segs.append((cid, ov_s, ov_e))\n    if not chunk_chain_segs:\n        chunk_chain_segs = [(\"A\", 0, len(chunk_seq))]\n    pdb_id   = \"VF\" + hashlib.md5(chunk_id.encode()).hexdigest()[:2].upper()\n    os.makedirs(mmcif_dir, exist_ok=True)\n    out_path = os.path.join(mmcif_dir, pdb_id + \".cif\")\n    _write_vfold_cif(chunk_seq, chunk_coords, chunk_chain_segs, pdb_id, out_path)\n    return out_path, pdb_id\n\n\ndef align_assembled_to_vfold(assembled, vfold_ref, min_anchors=10):\n    \"\"\"\n    Global Kabsch alignment of assembled Protenix output onto Vfold scaffold.\n\n    Eliminates the rotational/translational drift that accumulates when many\n    chunks are joined via stitch_chunks (each local alignment inherits error\n    from the previous). The Vfold scaffold provides an independent global\n    reference frame for every chunk simultaneously.\n\n    Returns aligned array, or the original if fewer than min_anchors anchors.\n    \"\"\"\n    if assembled.shape != vfold_ref.shape or len(assembled) < min_anchors:\n        return assembled\n    valid = (\n        (np.abs(assembled).sum(axis=1) > 1e-6) &\n        (np.abs(vfold_ref).sum(axis=1)  > 1e-6)\n    )\n    if valid.sum() < min_anchors:\n        return assembled\n    R, t = kabsch_alignment(assembled[valid], vfold_ref[valid])\n    return assembled @ R.T + t\n\n\n# ─────────────── Protenix Config Helpers ─────────────────────────────\ndef build_configs(input_json_path: str, dump_dir: str, model_name: str):\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(t, p):\n        for k, v in p.items():\n            if isinstance(v, dict) and k in t and isinstance(t[k], dict): deep_update(t[k], v)\n            else: t[k] = v\n\n    deep_update(base, model_configs[model_name])\n    mmcif_dir = os.environ.get(\"PROTENIX_MMCIF_DIR\", \"/kaggle/working/mmcif\")\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\"--data.template.prot_template_mmcif_dir {mmcif_dir}\",\n        f\"--sample_diffusion.N_sample {MODEL_N_SAMPLE}\",\n        f\"--seeds {SEED}\",\n        # ── Memory / precision tuning (previously missing) ──────────────────\n        # bf16 cuts VRAM usage ~50% vs fp32 with minimal accuracy impact.\n        f\"--dtype {DTYPE}\",\n        # Fewer Pairformer cycles = less intermediate activation memory.\n        f\"--model.N_cycle {N_CYCLE}\",\n        # Fewer diffusion steps = faster per-chunk inference on T4.\n        f\"--sample_diffusion.N_step {N_STEP}\",\n        # cuequivariance kernels use fused CUDA ops: lower VRAM + faster.\n        # NOTE: correct flag names are --triangle_multiplicative / --triangle_attention\n        f\"--triangle_multiplicative {TRIMUL_KERNEL}\",\n        f\"--triangle_attention {TRIATT_KERNEL}\",\n    ])\n    return parse_configs(configs=base, arg_str=arg_str, fill_required_with_null=True)\n\ndef coords_to_rows(target_id: str, seq: str, coords: np.ndarray) -> list:\n    rows = []\n    for i in range(len(seq)):\n        row = {\"ID\": f\"{target_id}_{i + 1}\", \"resname\": seq[i], \"resid\": i + 1}\n        for s in range(N_SAMPLE):\n            if s < coords.shape[0] and i < coords.shape[1]: x, y, z = coords[s, i]\n            else: x, y, z = 0.0, 0.0, 0.0\n            row[f\"x_{s + 1}\"] = float(x); row[f\"y_{s + 1}\"] = float(y); row[f\"z_{s + 1}\"] = float(z)\n        rows.append(row)\n    return rows\n\n# ─────────────── TBM Core Functions ──────────────────────────────────────────\ndef _make_aligner() -> PairwiseAligner:\n    al = PairwiseAligner()\n    al.mode = \"global\"\n    al.match_score = 2; al.mismatch_score = -1.5\n    al.open_gap_score = -8; al.extend_gap_score = -0.4\n    al.query_left_open_gap_score = -8; al.query_left_extend_gap_score = -0.4\n    al.query_right_open_gap_score = -8; al.query_right_extend_gap_score = -0.4\n    al.target_left_open_gap_score = -8; al.target_left_extend_gap_score = -0.4\n    al.target_right_open_gap_score = -8; al.target_right_extend_gap_score = -0.4\n    return al\n\n_aligner = _make_aligner()\n\ndef parse_stoichiometry(stoich: str) -> list:\n    if pd.isna(stoich) or str(stoich).strip() == \"\": return []\n    return [(ch.strip(), int(cnt)) for part in str(stoich).split(\";\") for ch, cnt in [part.split(\":\")]]\n\ndef parse_fasta(fasta_content: str) -> dict:\n    out, cur, parts = {}, None, []\n    for line in str(fasta_content).splitlines():\n        line = line.strip()\n        if not line: continue\n        if line.startswith(\">\"):\n            if cur is not None: out[cur] = \"\".join(parts)\n            cur = line[1:].split()[0]\n            parts = []\n        else: parts.append(line.replace(\" \", \"\"))\n    if cur is not None: out[cur] = \"\".join(parts)\n    return out\n\ndef get_chain_segments(row) -> list:\n    seq    = row[\"sequence\"]\n    stoich = row.get(\"stoichiometry\", \"\")\n    all_sq = row.get(\"all_sequences\", \"\")\n    if (pd.isna(stoich) or pd.isna(all_sq) or str(stoich).strip() == \"\" or str(all_sq).strip() == \"\"):\n        return [(0, len(seq))]\n    try:\n        chain_dict = parse_fasta(all_sq)\n        order = parse_stoichiometry(stoich)\n        segs, pos = [], 0\n        for ch, cnt in order:\n            base = chain_dict.get(ch)\n            if base is None: return [(0, len(seq))]\n            for _ in range(cnt):\n                segs.append((pos, pos + len(base)))\n                pos += len(base)\n        return segs if pos == len(seq) else [(0, len(seq))]\n    except Exception:\n        return [(0, len(seq))]\n\ndef parse_strands_from_row(row) -> list:\n    \"\"\"\n    Parse a CSV row into an ordered list of (chain_seq, count) pairs using\n    `all_sequences` (FASTA) + `stoichiometry` (e.g. \"A:2\" or \"B:1;A:1\").\n\n    Returns: [(chain_seq_str, copy_count), ...]  in stoichiometry order.\n    Consecutive identical chains are NOT merged so that order is preserved and\n    Protenix can correctly model asymmetric multimers.\n\n    Falls back to [(full_sequence, 1)] for monomers or when columns are absent.\n    \"\"\"\n    seq    = row[\"sequence\"]\n    stoich = row.get(\"stoichiometry\", \"\")\n    all_sq = row.get(\"all_sequences\",  \"\")\n\n    if pd.isna(stoich) or pd.isna(all_sq) or str(stoich).strip() == \"\" or str(all_sq).strip() == \"\":\n        return [(seq, 1)]\n    try:\n        chain_dict = parse_fasta(all_sq)\n        order      = parse_stoichiometry(stoich)\n        if not order:\n            return [(seq, 1)]\n\n        # Build list preserving stoichiometry order.\n        # Consecutive runs of the same chain_id with the same sequence get\n        # collapsed into a single entry with count = n_copies.\n        strands: list = []\n        for ch, cnt in order:\n            chain_seq = chain_dict.get(ch)\n            if chain_seq is None:\n                return [(seq, 1)]\n            # If the last entry has the same sequence, merge counts (homodimer opt.)\n            if strands and strands[-1][0] == chain_seq:\n                strands[-1] = (chain_seq, strands[-1][1] + cnt)\n            else:\n                strands.append((chain_seq, cnt))\n\n        # Sanity-check: reconstructed total length must match `sequence`\n        total = sum(len(s) * c for s, c in strands)\n        if total != len(seq):\n            return [(seq, 1)]\n\n        return strands\n    except Exception:\n        return [(seq, 1)]\n\n\ndef build_segments_map(df: pd.DataFrame) -> tuple:\n    seg_map, stoich_map = {}, {}\n    for _, r in df.iterrows():\n        tid = r[\"target_id\"]\n        seg_map[tid] = get_chain_segments(r)\n        stoich_map[tid] = \"\" if pd.isna(r.get(\"stoichiometry\", \"\")) else str(r.get(\"stoichiometry\", \"\"))\n    return seg_map, stoich_map\n\ndef process_labels(labels_df: pd.DataFrame) -> dict:\n    coords = {}\n    prefixes = labels_df[\"ID\"].str.rsplit(\"_\", n=1).str[0]\n    for prefix, grp in labels_df.groupby(prefixes):\n        coords[prefix] = grp.sort_values(\"resid\")[[\"x_1\", \"y_1\", \"z_1\"]].values\n    return coords\n\ndef _build_aligned_strings(query_seq, template_seq, alignment):\n    q_segs, t_segs = alignment.aligned\n    aq, at, qi, ti = [], [], 0, 0\n    for (qs, qe), (ts, te) in zip(q_segs, t_segs):\n        while qi < qs: aq.append(query_seq[qi]); at.append(\"-\"); qi += 1\n        while ti < ts: aq.append(\"-\"); at.append(template_seq[ti]); ti += 1\n        for qp, tp in zip(range(qs, qe), range(ts, te)):\n            aq.append(query_seq[qp]); at.append(template_seq[tp])\n        qi, ti = qe, te\n    while qi < len(query_seq): aq.append(query_seq[qi]); at.append(\"-\"); qi += 1\n    while ti < len(template_seq): aq.append(\"-\"); at.append(template_seq[ti]); ti += 1\n    return \"\".join(aq), \"\".join(at)\n\ndef find_similar_sequences_detailed(query_seq, query_ligand_fps, train_seqs_df, train_coords_dict, train_ligand_fps, quality_map=None, top_n=30):\n    \"\"\"TBM template search with ligand similarity, GC bonus, and structural quality ranking.\"\"\"\n    if quality_map is None:\n        quality_map = {}\n    results = []\n    for _, row in train_seqs_df.iterrows():\n        tid, tseq = row[\"target_id\"], row[\"sequence\"]\n        if tid not in train_coords_dict: continue\n        if abs(len(tseq) - len(query_seq)) / max(len(tseq), len(query_seq)) > 0.3: continue\n        \n        aln = next(iter(_aligner.align(query_seq, tseq)))\n        norm_s = aln.score / (2 * min(len(query_seq), len(tseq)))\n        \n        ligand_sim = 0.0\n        if query_ligand_fps:\n            t_fps = train_ligand_fps.get(tid, [])\n            ligand_sim = calc_max_ligand_similarity(query_ligand_fps, t_fps)\n            \n        gc_bonus = 0.0\n        identical = 0\n        for (qs, qe), (ts, te) in zip(*aln.aligned):\n            for qp, tp in zip(range(qs, qe), range(ts, te)):\n                if query_seq[qp] == tseq[tp]:\n                    identical += 1\n                    if query_seq[qp] in ['G', 'C']: gc_bonus += 1.0\n\n        # Quality boost from rna_metadata.csv (resolution + structuredness)\n        quality_bonus = 1.0\n        if tid in quality_map:\n            res, struct = quality_map[tid]\n            if res < 3.0:  # sub-3Å resolution = higher quality template\n                quality_bonus += (3.0 - res) * 0.05\n            if struct > 0.5:\n                quality_bonus += struct * 0.1\n\n        boosted_s = norm_s * (1.0 + (ligand_sim * 0.25)) * (1.0 + (gc_bonus / len(query_seq) * 0.15)) * quality_bonus\n        pct_id = 100 * identical / len(query_seq)\n        aq, at = _build_aligned_strings(query_seq, tseq, aln)\n        results.append((tid, tseq, boosted_s, train_coords_dict[tid], pct_id, aq, at))\n        \n    results.sort(key=lambda x: x[2], reverse=True)\n    return results[:top_n]\n\n# ─────────────── Innovation 2: Dynamic Fragment Re-threading ──────────────────\ndef identify_conserved_stems(query_seq: str, template_seq: str, alignment, pct_id: float) -> np.ndarray:\n    \"\"\"\n    When %identity < RETHREAD_IDENTITY_THRESHOLD, identify conserved stem regions\n    and flag loop regions for fragment re-threading (rather than rigid adaptation).\n    Returns a boolean mask: True = conserved stem, False = divergent loop to re-thread.\n    \"\"\"\n    conserved_mask = np.zeros(len(query_seq), dtype=bool)\n    \n    # Get base pairs to identify stems\n    try:\n        import RNA\n        fc = RNA.fold_compound(query_seq)\n        ss, _ = fc.mfe()\n        # Stems are base-paired positions in the MFE structure\n        in_stem = np.array([c in '()' for c in ss], dtype=bool)\n    except Exception:\n        in_stem = np.zeros(len(query_seq), dtype=bool)\n\n    # Mark aligned positions with high identity as conserved\n    for (qs, qe), (ts, te) in zip(*alignment.aligned):\n        for qp, tp in zip(range(qs, qe), range(ts, te)):\n            if query_seq[qp] == template_seq[tp]:\n                conserved_mask[qp] = True\n\n    # Conserved = aligned AND (in stem OR high overall identity)\n    if pct_id >= RETHREAD_IDENTITY_THRESHOLD:\n        return np.ones(len(query_seq), dtype=bool)  # All conserved at high identity\n    \n    return conserved_mask & in_stem\n\ndef adapt_template_to_query_rethreaded(query_seq, template_seq, template_coords,\n                                       pct_id: float, segments: list = None) -> tuple:\n    \"\"\"\n    Innovation 2: Dynamic Fragment Re-threading — now segment-aware for multimers.\n\n    - High identity (≥80%): rigid adaptation with linear gap fill.\n    - Low identity (<80%): torsion-space Catmull-Rom fill for loop gaps.\n    - Segment-aware: gap interpolation never crosses chain boundaries.\n      Each segment is filled independently using only anchors within that segment.\n    \"\"\"\n    if segments is None:\n        segments = [(0, len(query_seq))]\n\n    aln        = next(iter(_aligner.align(query_seq, template_seq)))\n    new_coords = np.full((len(query_seq), 3), np.nan)\n\n    for (qs, qe), (ts, te) in zip(*aln.aligned):\n        chunk = template_coords[ts:te]\n        if len(chunk) == (qe - qs):\n            new_coords[qs:qe] = chunk\n\n    valid_mask = ~np.isnan(new_coords[:, 0])\n\n    # Fill gaps segment-by-segment so inter-chain interpolation never happens\n    for seg_start, seg_end in segments:\n        for i in range(seg_start, seg_end):\n            if not np.isnan(new_coords[i, 0]):\n                continue\n\n            # Search for previous / next valid anchors WITHIN this segment only\n            pv = next((j for j in range(i - 1, seg_start - 1, -1)\n                       if not np.isnan(new_coords[j, 0])), -1)\n            nv = next((j for j in range(i + 1, seg_end)\n                       if not np.isnan(new_coords[j, 0])), -1)\n\n            if pv >= 0 and nv >= 0:\n                gap_len = nv - pv\n                t = (i - pv) / gap_len\n                if pct_id >= RETHREAD_IDENTITY_THRESHOLD:\n                    # Linear interpolation for high-identity templates\n                    new_coords[i] = (1 - t) * new_coords[pv] + t * new_coords[nv]\n                else:\n                    # Torsion-space Catmull-Rom path for low-identity loops\n                    linear = (1 - t) * new_coords[pv] + t * new_coords[nv]\n                    ang    = t * np.pi\n                    perp   = np.array([-np.sin(ang), np.cos(ang), 0.5 * np.sin(2 * ang)])\n                    perp  /= np.linalg.norm(perp) + 1e-9\n                    new_coords[i] = linear + 2.0 * (1.0 - pct_id / 100.0) * perp\n            elif pv >= 0:\n                new_coords[i] = new_coords[pv] + np.array([3, 0, 0])\n            elif nv >= 0:\n                new_coords[i] = new_coords[nv] + np.array([-3, 0, 0])\n            else:\n                # No anchors in this segment at all — place on A-form helix\n                new_coords[i] = np.array([\n                    9.0 * np.cos((i - seg_start) * 0.571),\n                    9.0 * np.sin((i - seg_start) * 0.571),\n                    (i - seg_start) * 2.81,\n                ])\n\n    return np.nan_to_num(new_coords), valid_mask\n\n\n# Keep monomer-compatible alias\ndef adapt_template_to_query(query_seq, template_seq, template_coords) -> tuple:\n    return adapt_template_to_query_rethreaded(query_seq, template_seq, template_coords, pct_id=100.0)\n\n# ─────────────── Innovation 1: Physics-Prior Diffusion Refinement ─────────────\n@njit(fastmath=True)\ndef physics_informed_refinement(coords, n, seg_starts, seg_ends, n_segs,\n                                ions_present=True, passes=200, strength=1.0):\n    \"\"\"\n    Segment-aware physics-informed backbone + tertiary-contact refiner.\n\n    Segment awareness (required for multimers):\n      - Backbone harmonic applies ONLY within a segment (no phantom covalent bond\n        between the tail of chain A and the head of chain B).\n      - Inter-chain docking potential: soft 6.5 Å pull for atom pairs from\n        DIFFERENT segments within a 10 Å influence sphere, to stabilise the\n        predicted interface without forcing unphysical contacts.\n\n    'strength' (0–1) scales all forces so high-confidence inputs (Protenix)\n    are not overwritten.\n    \"\"\"\n    backbone_force  = 0.08 * strength\n    intrachain_pull = (0.10 if ions_present else 0.02) * strength\n    interchain_pull = 0.04 * strength          # gentler inter-chain docking\n\n    # Precompute which segment each atom belongs to (O(n) once per call)\n    atom_seg = np.zeros(n, dtype=np.int64)\n    for s_idx in range(n_segs):\n        for a in range(seg_starts[s_idx], seg_ends[s_idx]):\n            atom_seg[a] = s_idx\n\n    # Precompute set of segment-boundary indices (last atom of each segment)\n    is_seg_end = np.zeros(n, dtype=np.bool_)\n    for s_idx in range(n_segs):\n        if seg_ends[s_idx] > 0:\n            is_seg_end[seg_ends[s_idx] - 1] = True\n\n    for _ in range(passes):\n        # 1. Backbone harmonic — skip across chain boundaries\n        for i in range(n - 1):\n            if is_seg_end[i]:\n                continue          # i is last residue of a chain; don't pull to i+1\n            vec_x = coords[i+1, 0] - coords[i, 0]\n            vec_y = coords[i+1, 1] - coords[i, 1]\n            vec_z = coords[i+1, 2] - coords[i, 2]\n            dist  = (vec_x**2 + vec_y**2 + vec_z**2) ** 0.5 + 1e-6\n            f     = (5.9 - dist) / dist * backbone_force\n            coords[i,   0] -= vec_x * f;  coords[i,   1] -= vec_y * f;  coords[i,   2] -= vec_z * f\n            coords[i+1, 0] += vec_x * f;  coords[i+1, 1] += vec_y * f;  coords[i+1, 2] += vec_z * f\n\n        # 2. Long-range potentials for pairs (i, j) with j ≥ i+15\n        if n > 20:\n            for i in range(n):\n                for j in range(i + 15, n):\n                    d_x  = coords[i, 0] - coords[j, 0]\n                    d_y  = coords[i, 1] - coords[j, 1]\n                    d_z  = coords[i, 2] - coords[j, 2]\n                    d_sq = d_x*d_x + d_y*d_y + d_z*d_z\n\n                    same_chain = (atom_seg[i] == atom_seg[j])\n\n                    if same_chain:\n                        # Intra-chain A-minor / tertiary contact — 8 Å target, 8 Å cutoff\n                        if d_sq < 64.0:\n                            d = d_sq ** 0.5 + 1e-6\n                            f = (8.0 - d) / d * intrachain_pull\n                            coords[i, 0] += d_x * f;  coords[i, 1] += d_y * f;  coords[i, 2] += d_z * f\n                            coords[j, 0] -= d_x * f;  coords[j, 1] -= d_y * f;  coords[j, 2] -= d_z * f\n                    else:\n                        # Inter-chain interface docking — 6.5 Å pull, 10 Å cutoff\n                        if d_sq < 100.0:\n                            d = d_sq ** 0.5 + 1e-6\n                            f = (6.5 - d) / d * interchain_pull\n                            coords[i, 0] += d_x * f;  coords[i, 1] += d_y * f;  coords[i, 2] += d_z * f\n                            coords[j, 0] -= d_x * f;  coords[j, 1] -= d_y * f;  coords[j, 2] -= d_z * f\n    return coords\n\n@njit(fastmath=True)\ndef apply_ion_stabilization_fast(coords, n, bridge_cutoff=10.0):\n    grad = np.zeros_like(coords)\n    bridge_strength = 0.15 \n    for i in range(n):\n        for j in range(i + 10, n): \n            dx = coords[i, 0] - coords[j, 0]; dy = coords[i, 1] - coords[j, 1]; dz = coords[i, 2] - coords[j, 2]\n            dist_sq = dx*dx + dy*dy + dz*dz\n            if dist_sq < bridge_cutoff*bridge_cutoff:\n                dist = np.sqrt(dist_sq) + 1e-6\n                force = (bridge_cutoff - dist) / dist\n                fx, fy, fz = dx*force*bridge_strength, dy*force*bridge_strength, dz*force*bridge_strength\n                grad[i, 0] -= fx; grad[i, 1] -= fy; grad[i, 2] -= fz\n                grad[j, 0] += fx; grad[j, 1] += fy; grad[j, 2] += fz\n    return grad\n\ndef adaptive_rna_constraints_with_ions(coords, target_id, segments_map, valid_mask=None,\n                                     ligand_pull_strength=0.0, confidence=1.0, passes=150,\n                                     ions_present=True):\n    \"\"\"\n    Backbone geometry corrector — segment-aware for multimers.\n\n    'confidence' controls how aggressively we override input coordinates:\n      - High (≥0.75, Protenix output): strength≈0.05; physics_informed_refinement\n        is NOT called. Protenix already satisfies geometry.\n      - Medium (≈0.7, TBM): strength≈0.23, sparse physics passes.\n      - Low (≈0.2, de-novo): strength≈0.60, regular physics passes.\n    \"\"\"\n    X        = coords.copy()\n    segments = segments_map.get(target_id, [(0, len(X))])\n    n_atoms  = len(X)\n    strength = max(0.75 * (1.0 - min(confidence, 0.97)), 0.05)\n\n    # Build numba-compatible segment arrays once (cheap, O(n_segs))\n    seg_starts_arr = np.array([s for s, e in segments], dtype=np.int64)\n    seg_ends_arr   = np.array([e for s, e in segments], dtype=np.int64)\n    n_segs         = len(segments)\n\n    apply_physics = (confidence < 0.75)\n\n    for k in range(passes):\n        # Per-segment backbone + self-clash correction\n        for s, e in segments:\n            C = X[s:e]; L = e - s\n            if L < 3: continue\n\n            d = C[1:] - C[:-1]; dist = np.linalg.norm(d, axis=1) + 1e-6\n            adj = d * ((5.95 - dist) / dist)[:, None] * (0.22 * strength)\n            C[:-1] -= adj; C[1:] += adj\n\n            if L > 2:\n                d2 = C[2:] - C[:-2]; d2n = np.linalg.norm(d2, axis=1) + 1e-6\n                adj2 = d2 * ((10.2 - d2n) / d2n)[:, None] * (0.10 * strength)\n                C[:-2] -= adj2; C[2:] += adj2\n\n            C[1:-1] += (0.06 * strength) * (0.5 * (C[:-2] + C[2:]) - C[1:-1])\n            if L >= 25:\n                idx  = np.linspace(0, L - 1, min(L, 160)).astype(int) if L > 220 else np.arange(L)\n                P    = C[idx]; diff = P[:, None, :] - P[None, :, :]\n                dm   = np.linalg.norm(diff, axis=2) + 1e-6\n                sep  = np.abs(idx[:, None] - idx[None, :])\n                mask = (sep > 2) & (dm < 3.2)\n                if np.any(mask):\n                    vec = (diff * ((3.2 - dm) / dm)[:, :, None] * mask[:, :, None]).sum(axis=1)\n                    C[idx] += (0.015 * strength) * vec\n            X[s:e] = C\n\n        # Physics passes: TBM / de-novo only, sparse, strength-scaled, segment-aware\n        if apply_physics and k > 60 and (k % 15 == 0):\n            X = physics_informed_refinement(\n                X, n_atoms, seg_starts_arr, seg_ends_arr, n_segs,\n                ions_present=ions_present, passes=3, strength=strength,\n            )\n\n        if ligand_pull_strength > 0 and valid_mask is not None:\n            gap_mask = ~valid_mask\n            if valid_mask.any() and gap_mask.any():\n                centroid = X[valid_mask].mean(axis=0)\n                diff     = centroid - X[gap_mask]\n                dists    = np.linalg.norm(diff, axis=1) + 1e-6\n                X[gap_mask] += (diff / dists[:, None]) * (ligand_pull_strength * 0.1 * strength)\n\n    return X\n\ndef _rotmat(axis, ang):\n    a = np.asarray(axis, float); a /= np.linalg.norm(a) + 1e-12\n    x, y, z = a; c, s = np.cos(ang), np.sin(ang); CC = 1 - c\n    return np.array([[c+x*x*CC, x*y*CC-z*s, x*z*CC+y*s],\n                     [y*x*CC+z*s, c+y*y*CC, y*z*CC-x*s],\n                     [z*x*CC-y*s, z*y*CC+x*s, c+z*z*CC]])\n\ndef apply_hinge(coords, seg, rng, deg=22):\n    s, e = seg; L = e - s\n    if L < 30: return coords\n    pivot = s + int(rng.integers(10, L - 10))\n    R = _rotmat(rng.normal(size=3), np.deg2rad(float(rng.uniform(-deg, deg))))\n    X = coords.copy(); p0 = X[pivot].copy()\n    X[pivot+1:e] = (X[pivot+1:e] - p0) @ R.T + p0\n    return X\n\ndef jitter_chains(coords, segs, rng, deg=12, trans=1.5):\n    X = coords.copy(); gc_ = X.mean(0, keepdims=True)\n    for s, e in segs:\n        R = _rotmat(rng.normal(size=3), np.deg2rad(float(rng.uniform(-deg, deg))))\n        shift = rng.normal(size=3); shift = shift / (np.linalg.norm(shift) + 1e-12) * float(rng.uniform(0, trans))\n        c = X[s:e].mean(0, keepdims=True)\n        X[s:e] = (X[s:e] - c) @ R.T + c + shift\n    X -= X.mean(0, keepdims=True) - gc_\n    return X\n\ndef smooth_wiggle(coords, segs, rng, amp=0.8):\n    X = coords.copy()\n    for s, e in segs:\n        L = e - s\n        if L < 20: continue\n        ctrl = np.linspace(0, L - 1, 6); disp = rng.normal(0, amp, (6, 3)); t = np.arange(L)\n        X[s:e] += np.vstack([np.interp(t, ctrl, disp[:, k]) for k in range(3)]).T\n    return X\n\n# ─────────────── Innovation 2: Torsion-Space Initial Structure ────────────────\ndef generate_rna_structure(sequence: str, seed=None) -> np.ndarray:\n    \"\"\"\n    Innovation 2: Replaces the unphysical spring helix with a torsion-angle\n    parameterized A-form RNA helix using standard η/θ backbone dihedrals.\n    This prevents knot formation and provides a valid starting geometry\n    for the physics refinement to converge from.\n    \n    Standard A-form RNA backbone torsion angles (mean values):\n      α=-68°, β=178°, γ=54°, δ=82°, ε=-153°, ζ=-71°, χ=-159°\n    The η/θ pseudo-torsion representation maps to a rise of ~2.8 Å/nt\n    and a twist of ~32.7°/nt in A-form helix.\n    \"\"\"\n    if seed is not None: np.random.seed(seed)\n    n = len(sequence)\n    coords = np.zeros((n, 3))\n    \n    # A-form helix parameters\n    rise_per_nt  = 2.81    # Å per nucleotide along helix axis\n    twist_per_nt = np.deg2rad(32.7)   # ~11 nt/turn in A-form\n    radius       = 9.0     # Å, radius of C1' from helix axis in A-form\n    \n    # Add small per-nucleotide torsion jitter to avoid perfect symmetry\n    # (improves convergence of the physics refinement)\n    jitter_scale = 0.3\n    \n    for i in range(n):\n        base = sequence[i]\n        # Purine (A, G) vs pyrimidine (C, U) have slightly different radii\n        r = radius + (0.3 if base in ('A', 'G') else -0.3)\n        ang = i * twist_per_nt + np.random.uniform(-jitter_scale, jitter_scale) * 0.1\n        z = i * rise_per_nt + np.random.uniform(-jitter_scale, jitter_scale) * 0.05\n        coords[i] = [r * np.cos(ang), r * np.sin(ang), z]\n    \n    return coords\n\n# ─────────────── Chain-Aware Dynamic Chunking ─────────────────────────────────\ndef generate_chain_aware_chunks(seq: str, segments: list, max_len: int = MAX_SEQ_LEN,\n                                overlap: int = CHUNK_OVERLAP) -> list:\n    \"\"\"\n    Generate chunks that respect RNA chain boundaries.\n    \n    Strategy:\n      1. If total sequence fits in one chunk → single chunk (no splitting needed).\n      2. If multiple chains exist, try to keep each chain in its own chunk.\n         Chains shorter than max_len are grouped together until the group\n         would exceed max_len.\n      3. Long chains (> max_len) are split with sliding window + overlap,\n         but the split points never cross into a different chain.\n    \n    Returns: list of (start, end) tuples, each representing a chunk window\n             into the concatenated sequence.\n    \"\"\"\n    L = len(seq)\n    \n    # Case 1: fits in one chunk\n    if L <= max_len:\n        return [(0, L)]\n    \n    # Case 2: single chain (monomer) — fall back to standard sliding window\n    if len(segments) <= 1:\n        chunks = []\n        pos = 0\n        while pos < L:\n            end = min(pos + max_len, L)\n            chunks.append((pos, end))\n            if end == L: break\n            pos = end - overlap\n        return chunks\n    \n    # Case 3: multi-chain — chain-aware grouping\n    chunks = []\n    group_start = segments[0][0]\n    group_end = segments[0][1]\n    \n    for seg_idx in range(1, len(segments)):\n        seg_s, seg_e = segments[seg_idx]\n        proposed_end = seg_e\n        proposed_len = proposed_end - group_start\n        \n        if proposed_len <= max_len:\n            # This chain fits in the current group\n            group_end = proposed_end\n        else:\n            # Flush current group — it may itself need sub-chunking\n            group_len = group_end - group_start\n            if group_len <= max_len:\n                chunks.append((group_start, group_end))\n            else:\n                # Sub-chunk this long group with sliding window\n                pos = group_start\n                while pos < group_end:\n                    end = min(pos + max_len, group_end)\n                    chunks.append((pos, end))\n                    if end == group_end: break\n                    pos = end - overlap\n            # Start new group with current segment\n            group_start = seg_s\n            group_end = seg_e\n    \n    # Flush final group\n    group_len = group_end - group_start\n    if group_len <= max_len:\n        chunks.append((group_start, group_end))\n    else:\n        pos = group_start\n        while pos < group_end:\n            end = min(pos + max_len, group_end)\n            chunks.append((pos, end))\n            if end == group_end: break\n            pos = end - overlap\n    \n    return chunks\n\n# ─────────────── Partner Molecule Extraction ──────────────────────────────────\ndef extract_partner_molecules(row) -> list:\n    \"\"\"\n    Parse all_sequences FASTA to find non-target molecules (proteins, DNA)\n    co-occurring in the experimental structure. Returns Protenix JSON\n    sequence entries for partner molecules.\n    \n    Protenix (AlphaFold3-based) can model RNA-protein and RNA-DNA complexes,\n    which dramatically improves accuracy for multi-molecular targets.\n    \n    Size guards: partners > MAX_PARTNER_RESIDUES are skipped to avoid OOM.\n    \"\"\"\n    stoich = row.get(\"stoichiometry\", \"\")\n    all_sq = row.get(\"all_sequences\", \"\")\n    if pd.isna(stoich) or pd.isna(all_sq) or str(stoich).strip() == \"\" or str(all_sq).strip() == \"\":\n        return []\n    \n    try:\n        chain_dict = parse_fasta(all_sq)\n    except Exception:\n        return []\n    \n    # Identify target chains IDs from stoichiometry\n    target_chains = set()\n    try:\n        for part in str(stoich).split(\";\"):\n            ch, _ = part.split(\":\")\n            target_chains.add(ch.strip())\n    except Exception:\n        return []\n    \n    partners = []\n    rna_chars = set(\"ACGU\")\n    dna_chars = set(\"ACGT\")\n    protein_chars = set(\"ACDEFGHIKLMNPQRSTVWY\")\n    \n    for header, seq_str in chain_dict.items():\n        chain_id = header.split()[0]\n        if chain_id in target_chains:\n            continue  # skip target RNA chains — already handled\n        \n        clean_seq = seq_str.upper().replace(\"-\", \"\").replace(\"N\", \"\").replace(\"X\", \"\")\n        if len(clean_seq) < 5:\n            continue  # skip very short fragments\n        if len(seq_str) > MAX_PARTNER_RESIDUES:\n            continue  # skip large partners to prevent OOM\n        \n        seq_chars = set(clean_seq)\n        \n        if seq_chars <= rna_chars:\n            partners.append({\"rnaSequence\": {\"sequence\": seq_str, \"count\": 1}})\n        elif seq_chars <= dna_chars:\n            partners.append({\"dnaSequence\": {\"sequence\": seq_str, \"count\": 1}})\n        elif seq_chars <= protein_chars:\n            partners.append({\"proteinChain\": {\"sequence\": seq_str, \"count\": 1}})\n    \n    return partners\n\n# ─────────────── Metadata Quality Map Builder ────────────────────────────────\ndef load_quality_map(metadata_path: str) -> dict:\n    \"\"\"Load rna_metadata.csv and build a quality lookup: pdb_id → (resolution, structuredness).\"\"\"\n    quality_map = {}\n    if not os.path.exists(metadata_path):\n        print(f\"  [Metadata] {metadata_path} not found — quality ranking disabled.\")\n        return quality_map\n    try:\n        meta_df = pd.read_csv(metadata_path)\n        for _, m in meta_df.iterrows():\n            pid = str(m.get(\"pdb_id\", m.get(\"entry_id\", \"\"))).strip().upper()\n            if not pid:\n                continue\n            res = m.get(\"resolution\", 99.0)\n            struct = m.get(\"total_adjusted_structuredness\", 0.0)\n            quality_map[pid] = (\n                float(res) if not pd.isna(res) else 99.0,\n                float(struct) if not pd.isna(struct) else 0.0,\n            )\n        print(f\"  [Metadata] Loaded quality data for {len(quality_map)} PDB entries\")\n    except Exception as e:\n        print(f\"  [Metadata] Failed to load: {e}\")\n    return quality_map\n\n\n# ─────────────── PDB_RNA Template Database (NEW: was defined but never used) ──\n\ndef load_pdb_seqres_db(fasta_path: str) -> dict:\n    \"\"\"\n    NEW: Load pdb_seqres_NA.fasta into a searchable dict {pdb_chain_key: sequence}.\n    Keys are like \"1ABC_A\". Only pure-RNA chains (ACGU) are kept.\n    This expands the TBM template pool from ~5K train/val entries to the full PDB.\n    \"\"\"\n    if not os.path.exists(fasta_path):\n        print(f\"  [PDB SEQRES] {fasta_path} not found — skipping PDB expansion\")\n        return {}\n    db = {}\n    rna_chars = set(\"ACGU\")\n    current_key, parts = None, []\n    try:\n        with open(fasta_path) as fh:\n            for line in fh:\n                line = line.strip()\n                if not line:\n                    continue\n                if line.startswith(\">\"):\n                    if current_key and parts:\n                        seq = \"\".join(parts).upper()\n                        if seq and set(seq) <= rna_chars and len(seq) >= 10:\n                            db[current_key] = seq\n                    current_key = line[1:].split()[0]\n                    parts = []\n                else:\n                    parts.append(line.replace(\" \", \"\").replace(\"-\", \"\").upper())\n        if current_key and parts:\n            seq = \"\".join(parts).upper()\n            if seq and set(seq) <= rna_chars and len(seq) >= 10:\n                db[current_key] = seq\n    except Exception as e:\n        print(f\"  [PDB SEQRES] Load failed: {e}\")\n    print(f\"  [PDB SEQRES] Loaded {len(db):,} RNA chains from PDB\")\n    return db\n\n\ndef load_pdb_release_dates(csv_path: str) -> dict:\n    \"\"\"NEW: Load pdb_release_dates_NA.csv → {PDB_ID_UPPER: 'YYYY-MM-DD'} for temporal filtering.\"\"\"\n    if not os.path.exists(csv_path):\n        return {}\n    try:\n        df = pd.read_csv(csv_path)\n        id_col   = next((c for c in df.columns if any(x in c.lower() for x in (\"id\", \"entry\", \"pdb\"))), df.columns[0])\n        date_col = next((c for c in df.columns if any(x in c.lower() for x in (\"date\", \"release\"))), df.columns[1])\n        result = {str(r[id_col]).strip().upper(): str(r[date_col]) for _, r in df.iterrows()}\n        print(f\"  [Release Dates] Loaded {len(result):,} PDB release dates\")\n        return result\n    except Exception as e:\n        print(f\"  [Release Dates] Failed: {e}\")\n        return {}\n\n\ndef extract_c1prime_from_cif(pdb_id: str, cif_dir: str) -> np.ndarray:\n    \"\"\"\n    NEW: Parse an mmCIF file and return C1' atom coordinates for all RNA residues.\n    Returns shape (n_rna_residues, 3) or None on failure.\n    Only model 1, only canonical RNA residues (A/C/G/U).\n    Uses a streaming line-by-line parser to handle large CIF files efficiently.\n    \"\"\"\n    for suffix in (pdb_id.lower(), pdb_id.upper()):\n        path = os.path.join(cif_dir, f\"{suffix}.cif\")\n        if os.path.exists(path):\n            break\n    else:\n        return None\n\n    try:\n        columns, rows, in_loop, reading_atom = [], [], False, False\n        rna_res = {\"A\", \"C\", \"G\", \"U\"}\n\n        with open(path, \"r\") as fh:\n            for raw in fh:\n                line = raw.rstrip()\n                stripped = line.strip()\n\n                if stripped == \"loop_\":\n                    # Flush previous atom site block if any\n                    in_loop, reading_atom, columns = True, False, []\n                    continue\n\n                if in_loop and stripped.startswith(\"_atom_site.\"):\n                    columns.append(stripped.split(\".\", 1)[1].split()[0])\n                    reading_atom = True\n                    continue\n\n                if reading_atom and columns:\n                    if stripped.startswith(\"_\") or stripped == \"loop_\" or stripped.startswith(\"#\"):\n                        reading_atom = False\n                        in_loop = stripped == \"loop_\"\n                        if in_loop:\n                            columns = []\n                        continue\n                    if stripped and not stripped.startswith(\"#\"):\n                        parts = stripped.split()\n                        if len(parts) >= len(columns):\n                            rows.append(parts[: len(columns)])\n\n        if not columns or not rows:\n            return None\n\n        ci = {c: i for i, c in enumerate(columns)}\n\n        # Resolve column names (label_ preferred, auth_ fallback)\n        def col(*names):\n            for n in names:\n                if n in ci:\n                    return ci[n]\n            return -1\n\n        g_col  = col(\"group_PDB\")\n        at_col = col(\"label_atom_id\", \"auth_atom_id\")\n        cp_col = col(\"label_comp_id\", \"auth_comp_id\")\n        ch_col = col(\"label_asym_id\", \"auth_asym_id\")\n        rs_col = col(\"label_seq_id\",  \"auth_seq_id\")\n        x_col  = col(\"Cartn_x\")\n        y_col  = col(\"Cartn_y\")\n        z_col  = col(\"Cartn_z\")\n        m_col  = col(\"pdbx_PDB_model_num\")\n\n        if -1 in (at_col, cp_col, ch_col, rs_col, x_col, y_col, z_col):\n            return None\n\n        atoms = []\n        for row in rows:\n            try:\n                if g_col >= 0 and row[g_col] != \"ATOM\":\n                    continue\n                if m_col >= 0 and row[m_col] not in (\"1\", \".\"):\n                    continue\n                atom_id = row[at_col].strip(\"'\\\"\")\n                if atom_id not in (\"C1'\", \"C1*\"):\n                    continue\n                comp = row[cp_col]\n                if comp not in rna_res:\n                    continue\n                resid_str = row[rs_col]\n                resid = int(resid_str) if resid_str not in (\".\", \"?\") else 0\n                atoms.append((row[ch_col], resid, float(row[x_col]), float(row[y_col]), float(row[z_col])))\n            except (ValueError, IndexError):\n                continue\n\n        if not atoms:\n            return None\n        atoms.sort(key=lambda a: (a[0], a[1]))\n        return np.array([[a[2], a[3], a[4]] for a in atoms], dtype=np.float32)\n\n    except Exception:\n        return None\n\n\n# LRU-style CIF coordinate cache (avoids re-parsing the same file)\n_cif_coord_cache: dict = {}\n_CIF_CACHE_LIMIT = 600\n\n\ndef get_cif_coords(pdb_id: str, cif_dir: str) -> np.ndarray:\n    \"\"\"NEW: Cached wrapper around extract_c1prime_from_cif().\"\"\"\n    pid = pdb_id.upper()\n    if pid not in _cif_coord_cache:\n        coords = extract_c1prime_from_cif(pid, cif_dir)\n        if len(_cif_coord_cache) < _CIF_CACHE_LIMIT:\n            _cif_coord_cache[pid] = coords\n        return coords\n    return _cif_coord_cache[pid]\n\n\ndef build_kmer_index(seq_dict: dict, k: int = 5) -> dict:\n    \"\"\"NEW: Build an inverted k-mer index for fast candidate pre-filtering.\"\"\"\n    index: dict = {}\n    for key, seq in seq_dict.items():\n        for i in range(len(seq) - k + 1):\n            kmer = seq[i : i + k]\n            if kmer not in index:\n                index[kmer] = []\n            index[kmer].append(key)\n    return index\n\n\ndef kmer_candidates(query_seq: str, kmer_index: dict, k: int = 5, min_hits: int = 4) -> set:\n    \"\"\"NEW: Return keys sharing at least min_hits k-mers with query_seq.\"\"\"\n    counts: dict = {}\n    for i in range(len(query_seq) - k + 1):\n        kmer = query_seq[i : i + k]\n        if kmer in kmer_index:\n            for key in kmer_index[kmer]:\n                counts[key] = counts.get(key, 0) + 1\n    return {k for k, v in counts.items() if v >= min_hits}\n\n\ndef parse_msa_for_pdb_hits(target_id: str, msa_dir: str) -> list:\n    \"\"\"\n    NEW: Parse {target_id}.MSA.fasta and extract (pdb_id, chain_id, sequence) tuples\n    for all homologs that can be traced back to a PDB entry.\n    These are *directly known* homologs — far more reliable starting points for\n    template search than a blind SEQRES scan.\n    \"\"\"\n    msa_path = os.path.join(msa_dir, f\"{target_id}.MSA.fasta\")\n    if not os.path.exists(msa_path):\n        return []\n    hits = []\n    # PDB IDs: 4-char alphanumeric starting with a digit, e.g. \"1ABC\"\n    pdb_re    = re.compile(r\"(?i)(?:pdb[|_:])?([0-9][A-Za-z0-9]{3})[|_]([A-Za-z0-9]+)\")\n    chain_re  = re.compile(r\"chain=([A-Za-z0-9]+)\")\n    try:\n        cur_hdr, cur_parts = None, []\n        with open(msa_path) as fh:\n            for line in fh:\n                line = line.strip()\n                if not line:\n                    continue\n                if line.startswith(\">\"):\n                    if cur_hdr and cur_parts:\n                        seq = \"\".join(cur_parts).upper().replace(\"-\", \"\")\n                        m = pdb_re.search(cur_hdr)\n                        if m and seq:\n                            pdb_id   = m.group(1).upper()\n                            chain_id = chain_re.search(cur_hdr)\n                            chain_id = chain_id.group(1) if chain_id else m.group(2)\n                            hits.append((pdb_id, chain_id, seq))\n                    cur_hdr, cur_parts = line[1:], []\n                else:\n                    cur_parts.append(line)\n        if cur_hdr and cur_parts:\n            seq = \"\".join(cur_parts).upper().replace(\"-\", \"\")\n            m = pdb_re.search(cur_hdr)\n            if m and seq:\n                pdb_id   = m.group(1).upper()\n                chain_id = chain_re.search(cur_hdr)\n                chain_id = chain_id.group(1) if chain_id else m.group(2)\n                hits.append((pdb_id, chain_id, seq))\n    except Exception:\n        pass\n    return hits\n\n\ndef search_pdb_seqres_for_templates(\n    query_seq: str,\n    query_ligand_fps: list,\n    pdb_seqres_dict: dict,\n    pdb_kmer_index: dict,\n    cif_dir: str,\n    release_dates: dict,\n    quality_map: dict,\n    training_cutoff: str = \"2025-05-29\",\n    top_n: int = 20,\n) -> list:\n    \"\"\"\n    NEW: Search pdb_seqres_NA.fasta for templates beyond the train/val set.\n    Two-stage:\n      1. k-mer pre-filter (fast): candidates sharing ≥4 5-mers with query\n      2. Pairwise alignment (slow): only on pre-filtered candidates\n    Templates released AFTER training_cutoff are excluded to avoid leakage.\n    Returns list of (pdb_chain_key, seq, boosted_score, coords, pct_id, aq, at).\n    \"\"\"\n    candidates = kmer_candidates(query_seq, pdb_kmer_index, k=5, min_hits=4)\n    # Also apply length filter to drop hopeless candidates fast\n    q_len = len(query_seq)\n    candidates = {\n        k for k in candidates\n        if abs(len(pdb_seqres_dict[k]) - q_len) / max(len(pdb_seqres_dict[k]), q_len) <= 0.35\n    }\n\n    results = []\n    for key in candidates:\n        tseq = pdb_seqres_dict[key]\n        # Temporal filter: skip structures released after training cutoff\n        pdb_id = key.split(\"_\")[0].upper()\n        rel_date = release_dates.get(pdb_id, \"\")\n        if rel_date and rel_date > training_cutoff:\n            continue\n\n        coords = get_cif_coords(pdb_id, cif_dir)\n        if coords is None or len(coords) < 5:\n            continue\n\n        # Coords length may differ from SEQRES length (missing residues, etc.)\n        # Use the coords directly as the \"template\" up to min(len(tseq), len(coords))\n        usable_len = min(len(tseq), len(coords))\n        if usable_len < 5:\n            continue\n        tseq_used  = tseq[:usable_len]\n        coords_used = coords[:usable_len]\n\n        aln    = next(iter(_aligner.align(query_seq, tseq_used)))\n        norm_s = aln.score / (2 * min(len(query_seq), len(tseq_used)))\n        if norm_s < MIN_SIMILARITY:\n            continue\n\n        identical = 0\n        gc_bonus  = 0.0\n        for (qs, qe), (ts, te) in zip(*aln.aligned):\n            for qp, tp in zip(range(qs, qe), range(ts, te)):\n                if query_seq[qp] == tseq_used[tp]:\n                    identical += 1\n                    if query_seq[qp] in (\"G\", \"C\"):\n                        gc_bonus += 1.0\n\n        pct_id   = 100.0 * identical / q_len\n        if pct_id < MIN_PERCENT_IDENTITY:\n            continue\n\n        quality_bonus = 1.0\n        if pdb_id in quality_map:\n            res, struct = quality_map[pdb_id]\n            if res < 3.0:\n                quality_bonus += (3.0 - res) * 0.05\n            if struct > 0.5:\n                quality_bonus += struct * 0.1\n\n        ligand_sim = 0.0\n        if query_ligand_fps and RDKIT_AVAILABLE:\n            ligand_sim = 0.0  # No SMILES for PDB SEQRES entries; leave at 0\n\n        boosted = (\n            norm_s\n            * (1.0 + ligand_sim * 0.25)\n            * (1.0 + gc_bonus / q_len * 0.15)\n            * quality_bonus\n        )\n\n        aq, at = _build_aligned_strings(query_seq, tseq_used, aln)\n        results.append((key, tseq_used, boosted, coords_used, pct_id, aq, at))\n\n    results.sort(key=lambda x: x[2], reverse=True)\n    return results[:top_n]\n\n# ─────────────── TBM Phase ───────────────────────────────────────────────────\ndef tbm_phase(test_df, train_seqs_df, train_coords_dict, segments_map, train_ligand_fps, quality_map=None, pdb_seqres_dict=None, pdb_kmer_index=None, release_dates=None):\n    print(f\"\\n{'='*60}\\nPHASE 1: Template-Based Modeling (Ligand-Boosted + Re-threading)\\n{'='*60}\")\n    t0 = time.time()\n\n    template_predictions: dict = {}\n    protenix_queue:       dict = {}\n\n    for _, row in test_df.iterrows():\n        tid, seq = row[\"target_id\"], row[\"sequence\"]\n        segs = segments_map.get(tid, [(0, len(seq))])\n        \n        q_smiles = row.get(\"ligand_SMILES\", \"\")\n        q_ids    = row.get(\"ligand_ids\",    \"\")\n        q_fps, max_heavy_atoms = extract_ligand_features(q_smiles, q_ids)\n\n        # ions_present: True if any KNOWN_ION_CCD_CODES appear in the ligand list.\n        # Used by physics_informed_refinement to scale the Debye-Hückel pull factor.\n        ccd_codes_present = {c.strip().upper() for c in str(q_ids).split(\";\") if c.strip()}\n        ions_present_flag = bool(ccd_codes_present & KNOWN_ION_CCD_CODES)\n\n        # Innovation 3: Scale ligand pull strength by organic heavy-atom count only\n        ligand_gravity = min(max_heavy_atoms / 40.0, 1.0) if max_heavy_atoms > 0 else 0.0\n\n        # Search train/val templates (original)\n        similar = find_similar_sequences_detailed(seq, q_fps, train_seqs_df, train_coords_dict, train_ligand_fps, quality_map=quality_map, top_n=30)\n\n        # NEW: Augment with MSA-guided PDB hits (direct homologs from MSA files)\n        msa_hits_used = set()\n        if os.path.isdir(MSA_DIR):\n            msa_pdb_hits = parse_msa_for_pdb_hits(tid, MSA_DIR)\n            for pdb_id, chain_id, msa_seq in msa_pdb_hits[:40]:\n                key = f\"{pdb_id}_{chain_id}\"\n                if key in msa_hits_used:\n                    continue\n                msa_coords = get_cif_coords(pdb_id, PDB_RNA_CIF_DIR)\n                if msa_coords is None or len(msa_coords) < 5:\n                    continue\n                usable = min(len(msa_seq), len(msa_coords))\n                if usable < 5:\n                    continue\n                aln = next(iter(_aligner.align(seq, msa_seq[:usable])))\n                norm_s = aln.score / (2 * min(len(seq), usable))\n                identical = sum(\n                    1 for (qs, qe), (ts, te) in zip(*aln.aligned)\n                    for qp, tp in zip(range(qs, qe), range(ts, te))\n                    if seq[qp] == msa_seq[tp]\n                )\n                pct_id_msa = 100.0 * identical / len(seq)\n                if norm_s >= MIN_SIMILARITY and pct_id_msa >= MIN_PERCENT_IDENTITY:\n                    aq, at = _build_aligned_strings(seq, msa_seq[:usable], aln)\n                    # Prepend MSA hits so they get priority slots\n                    similar.insert(0, (key, msa_seq[:usable], norm_s * 1.05, msa_coords[:usable], pct_id_msa, aq, at))\n                    msa_hits_used.add(key)\n\n        # NEW: Search expanded PDB SEQRES DB for additional templates\n        if pdb_seqres_dict and pdb_kmer_index and len(similar) < N_SAMPLE:\n            pdb_extra = search_pdb_seqres_for_templates(\n                seq, q_fps, pdb_seqres_dict, pdb_kmer_index,\n                PDB_RNA_CIF_DIR, release_dates or {}, quality_map or {},\n                top_n=20,\n            )\n            # Merge: avoid duplicates with existing similar list\n            existing_ids = {r[0] for r in similar}\n            for entry in pdb_extra:\n                if entry[0] not in existing_ids:\n                    similar.append(entry)\n            similar.sort(key=lambda x: x[2], reverse=True)\n        preds, used = [], set()\n\n        for i, (tmpl_id, tmpl_seq, sim, tmpl_coords, pct_id, _, _) in enumerate(similar):\n            if len(preds) >= N_SAMPLE - 1: break\n            if sim < MIN_SIMILARITY or pct_id < MIN_PERCENT_IDENTITY: break           \n            if tmpl_id in used: continue\n\n            rng = np.random.default_rng((row.name * 10000000000 + i * 10007) % (2**32))\n            \n            # Innovation 2: pass segments so gap fill never crosses chain boundaries\n            adapted, valid_mask = adapt_template_to_query_rethreaded(\n                seq, tmpl_seq, tmpl_coords, pct_id, segments=segs\n            )\n\n            slot = len(preds)\n            if slot == 0: X = adapted\n            elif slot == 1: X = adapted + rng.normal(0, max(0.01, (0.40 - sim) * 0.06), adapted.shape)\n            elif slot == 2: longest = max(segs, key=lambda se: se[1] - se[0]); X = apply_hinge(adapted, longest, rng)\n            elif slot == 3: X = jitter_chains(adapted, segs, rng)\n            else: X = smooth_wiggle(adapted, segs, rng)\n\n            refined = adaptive_rna_constraints_with_ions(\n                X, tid, segments_map, valid_mask,\n                ligand_pull_strength=ligand_gravity,\n                confidence=sim, passes=120,\n                ions_present=ions_present_flag,\n            )\n            preds.append(refined)\n            used.add(tmpl_id)\n\n        template_predictions[tid] = preds\n        n_needed = N_SAMPLE - len(preds)\n        \n        if n_needed > 0:\n            # Chain-aware dynamic chunking — never splits across chain boundaries\n            chunks_info = generate_chain_aware_chunks(seq, segs, MAX_SEQ_LEN, CHUNK_OVERLAP)\n                \n            protenix_queue[tid] = {\n                \"n_needed\":       n_needed,\n                \"seq\":            seq,\n                \"row\":            row,        # keep row for partner molecule extraction\n                \"strands\":        parse_strands_from_row(row),   # [(chain_seq, count), ...]\n                \"segments\":       segs,                           # [(start, end), ...]\n                \"chunks\":         chunks_info,\n                \"max_heavy_atoms\": max_heavy_atoms,\n                \"ions_present\":   ions_present_flag,\n            }\n            print(f\"  {tid} ({len(seq)} nt): {len(preds)} TBM → Queued {len(chunks_info)} chunks for Protenix\")\n        else:\n            print(f\"  {tid} ({len(seq)} nt): all {N_SAMPLE} from TBM ✓\")\n\n    print(f\"\\nPhase 1 done in {time.time() - t0:.1f}s\")\n    return template_predictions, protenix_queue\n\n# ─────────────── Ligand / Ion JSON Entity Classifier ────────────────────────\ndef classify_ligand_entities(row) -> list:\n    \"\"\"\n    Parse a row's ligand_ids and ligand_SMILES and return a list of correctly-typed\n    Protenix JSON sequence entries — either {\"ion\": ...} or {\"ligand\": ...}.\n\n    The json_parser.py routing rules are:\n      {\"ion\":    {\"ion\":    \"MG\",       \"count\": N}} → CCD lookup, NO EmbedMolecule\n      {\"ligand\": {\"ligand\": \"CCD_ATP\",  \"count\": N}} → CCD lookup, NO EmbedMolecule\n      {\"ligand\": {\"ligand\": \"CC(=O)O\",  \"count\": N}} → smiles_to_atom_info → EmbedMolecule\n\n    Strategy (in priority order):\n      1. If the CCD code is in KNOWN_ION_CCD_CODES → use \"ion\" key.\n      2. If we have a CCD code (even for an organic ligand) → use \"ligand\" with \"CCD_XXX\"\n         to avoid EmbedMolecule entirely.\n      3. Fall back to SMILES only for novel organic ligands with no CCD code, after\n         confirming RDKit can generate a valid 3D conformer.\n    \"\"\"\n    entities = []\n\n    raw_ids   = str(row.get(\"ligand_ids\",    \"\") or \"\").strip()\n    raw_smis  = str(row.get(\"ligand_SMILES\", \"\") or \"\").strip()\n\n    if not raw_smis:\n        return entities\n\n    smi_list = [s.strip() for s in raw_smis.split(\";\") if s.strip()]\n    ccd_list = [c.strip().upper() for c in raw_ids.split(\";\")] if raw_ids else []\n\n    for idx, smi in enumerate(smi_list):\n        ccd_code = ccd_list[idx] if idx < len(ccd_list) else \"\"\n\n        # ── Path 1: Known ion → use \"ion\" key with CCD code (zero RDKit involvement)\n        if ccd_code and ccd_code in KNOWN_ION_CCD_CODES:\n            entities.append({\"ion\": {\"ion\": ccd_code, \"count\": 1}})\n            continue\n\n        # ── Path 2: Known organic ligand with a CCD code → \"CCD_XXX\" avoids EmbedMolecule\n        if ccd_code:\n            entities.append({\"ligand\": {\"ligand\": f\"CCD_{ccd_code}\", \"count\": 1}})\n            continue\n\n        # ── Path 3: No CCD code — must use SMILES. Validate with RDKit first.\n        if not RDKIT_AVAILABLE:\n            continue\n        try:\n            mol = Chem.MolFromSmiles(smi)\n            if mol is None or mol.GetNumHeavyAtoms() < 3:\n                continue\n            # Quick single-atom ion check even without CCD code\n            if mol.GetNumHeavyAtoms() == 1:\n                continue\n            mol_h = Chem.AddHs(mol)\n            ec = AllChem.EmbedMolecule(mol_h)\n            if ec != 0:\n                ec = AllChem.EmbedMolecule(mol_h, useRandomCoords=True)\n            if ec != 0:\n                print(f\"  [Ligand Skip] EmbedMolecule failed for SMILES, dropping: {smi[:60]}\")\n                continue\n            canonical = Chem.MolToSmiles(mol, isomericSmiles=True)\n            entities.append({\"ligand\": {\"ligand\": canonical, \"count\": 1}})\n        except Exception as e:\n            print(f\"  [Ligand Skip] Exception processing SMILES '{smi[:40]}': {e}\")\n            continue\n\n    return entities\n\n# ─────────────── Chemical-Guided Chunk JSON Builder ─────────────────────\ndef build_input_json_chunked(queue_dict: dict, df: pd.DataFrame, json_path: str) -> None:\n    \"\"\"\n    Build Protenix input JSON with:\n      1. True multimer support — each chain is a separate rnaSequence entry so\n         Protenix models the docking interface without phantom covalent bonds.\n      2. Correct ion/ligand entity typing (ion → CCD path, no EmbedMolecule).\n\n    For single-chunk entries (total length ≤ MAX_SEQ_LEN):\n      Each unique chain gets its own {\"rnaSequence\": {\"sequence\": ..., \"count\": N}}.\n\n    For multi-chunk entries (long sequences):\n      Each chunk is mapped back to its chain segments so no cross-chain concat\n      is ever sent as one rnaSequence. Separate rnaSequence entries are emitted\n      for each chain portion within the chunk window.\n\n    json_parser.py routing rules:\n      {\"ion\":    {\"ion\":    \"MG\",      \"count\":1}} → CCD, no EmbedMolecule\n      {\"ligand\": {\"ligand\": \"CCD_ATP\", \"count\":1}} → CCD, no EmbedMolecule\n      {\"ligand\": {\"ligand\": \"CC(=O)O\", \"count\":1}} → smiles → EmbedMolecule\n    \"\"\"\n    data = []\n\n    for tid, info in queue_dict.items():\n        full_seq        = info[\"seq\"]\n        strands         = info[\"strands\"]          # [(chain_seq, count), ...]\n        segments        = info[\"segments\"]          # [(start, end), ...]  per-chain instance\n        row             = df[df[\"target_id\"] == tid].iloc[0]\n        ligand_entities = classify_ligand_entities(row)\n\n        is_single_chunk = (len(info[\"chunks\"]) == 1)\n\n        for c_idx, (start, end) in enumerate(info[\"chunks\"]):\n            chunk_id = f\"{tid}_chunk{c_idx}\"\n\n            entry = {\n                \"name\":           chunk_id,\n                \"covalent_bonds\": [],\n                \"sequences\":      [],\n            }\n\n            if is_single_chunk:\n                # ── True multimer: one rnaSequence per unique chain with correct count ──\n                for chain_seq, chain_count in strands:\n                    entry[\"sequences\"].append({\n                        \"rnaSequence\": {\"sequence\": chain_seq, \"count\": chain_count}\n                    })\n            else:\n                # ── Multi-chunk: split chunk at chain boundaries ──\n                # For each chain segment that overlaps this chunk window,\n                # emit only the overlapping portion as a separate rnaSequence.\n                # This avoids a single rnaSequence spanning two different chains.\n                added_any = False\n                for seg_start, seg_end in segments:\n                    overlap_s = max(start, seg_start)\n                    overlap_e = min(end,   seg_end)\n                    if overlap_s >= overlap_e:\n                        continue\n                    chain_portion = full_seq[overlap_s:overlap_e]\n                    entry[\"sequences\"].append({\n                        \"rnaSequence\": {\"sequence\": chain_portion, \"count\": 1}\n                    })\n                    added_any = True\n\n                # Safety fallback: if segment map is empty/wrong, use raw slice\n                if not added_any:\n                    entry[\"sequences\"].append({\n                        \"rnaSequence\": {\"sequence\": full_seq[start:end], \"count\": 1}\n                    })\n\n            # Add partner molecules (protein/DNA) for co-folding\n            partner_entities = extract_partner_molecules(row)\n            # Size guard: check total token count before adding partners\n            rna_tokens = sum(len(s.get(\"rnaSequence\", {}).get(\"sequence\", \"\"))\n                             * s.get(\"rnaSequence\", {}).get(\"count\", 1)\n                             for s in entry[\"sequences\"] if \"rnaSequence\" in s)\n            partner_tokens = sum(\n                len(list(p.values())[0].get(\"sequence\", \"\")) for p in partner_entities\n            )\n            if rna_tokens + partner_tokens <= MAX_TOTAL_TOKENS:\n                entry[\"sequences\"].extend(partner_entities)\n                if partner_entities:\n                    print(f\"    {chunk_id}: +{len(partner_entities)} partner molecule(s) for co-folding\")\n            elif partner_entities:\n                print(f\"    {chunk_id}: skipped {len(partner_entities)} partners (would exceed {MAX_TOTAL_TOKENS} tokens)\")\n\n            # Append ligand / ion entities with correct JSON keys\n            entry[\"sequences\"].extend(ligand_entities)\n\n            data.append(entry)\n\n    with open(json_path, \"w\", encoding=\"utf-8\") as f:\n        json.dump(data, f)\n\n# ─────────────── Multi-GPU Helpers ───────────────────────────────────────────\n\ndef split_queue_across_gpus(queue: dict, n_gpus: int = 2) -> list:\n    \"\"\"\n    Greedy split of protenix_queue across n_gpus workers based on sequence length.\n    This balances the nucleotide load (and thus compute time) roughly equally across\n    available GPUs, preventing one GPU from finishing long before the other.\n    \"\"\"\n    # Sort targets by sequence length descending\n    items = sorted(queue.items(), key=lambda x: len(x[1][\"seq\"]), reverse=True)\n    \n    splits: list = [{} for _ in range(n_gpus)]\n    gpu_loads = [0] * n_gpus\n    \n    for tid, info in items:\n        # Find the GPU that currently has the lightest load\n        min_load_idx = gpu_loads.index(min(gpu_loads))\n        # Assign target to this GPU\n        splits[min_load_idx][tid] = info\n        gpu_loads[min_load_idx] += len(info[\"seq\"])\n        \n    print(f\"  Queue split loads (in total nucleotides): {gpu_loads}\")\n    return splits\n\n\ndef write_gpu_worker_script(script_path: str) -> None:\n    \"\"\"\n    Write a self-contained Protenix single-GPU inference worker to disk.\n    Contains strictly inline dependencies so the child process works seamlessly.\n    \"\"\"\n    script = r'''#!/usr/bin/env python3\nimport argparse, os, sys, gc, json, pickle, warnings\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom numba import njit\nfrom tqdm import tqdm\n\nwarnings.filterwarnings(\"ignore\")\n\ntry:\n    import RNA as _RNA_LIB\n    _VIENNA_AVAILABLE = True\nexcept ImportError:\n    _VIENNA_AVAILABLE = False\n\ntry:\n    from rdkit import Chem\n    from rdkit.Chem import AllChem\n    _RDKIT_AVAILABLE = True\nexcept ImportError:\n    _RDKIT_AVAILABLE = False\n\n# Use CoDock as configured in main parent script\nUSE_CODOCK = True\n\ndef compute_delta_g(seq: str) -> float:\n    if not _VIENNA_AVAILABLE or not seq:\n        return 0.0\n    try:\n        fc = _RNA_LIB.fold_compound(seq)\n        _, mfe = fc.mfe()\n        return float(mfe)\n    except Exception:\n        return 0.0\n\ndef codock_score_sample(rna_c1_coords: np.ndarray, smiles_list: list) -> float:\n    if not _RDKIT_AVAILABLE or not smiles_list:\n        return 0.0\n\n    N_CONF         = 10\n    CONTACT_CUTOFF = 15.0\n    CLASH_THR      = 2.5\n    CLASH_PENALTY  = 5.0\n\n    pocket_centre = rna_c1_coords.mean(axis=0)\n    rna_arr = rna_c1_coords.astype(np.float32)\n    total_score = 0.0\n\n    for smi in smiles_list:\n        if not isinstance(smi, str):\n            continue\n        smi = smi.strip()\n        if not smi or smi.lower() == \"nan\":\n            continue\n        mol = Chem.MolFromSmiles(smi)\n        if mol is None or mol.GetNumHeavyAtoms() < 3:\n            continue\n        mol_h = Chem.AddHs(mol)\n\n        params = AllChem.ETKDGv3()\n        params.randomSeed = 42\n        params.numThreads = 1\n        ids = AllChem.EmbedMultipleConfs(mol_h, numConfs=N_CONF, params=params)\n        if len(ids) == 0:\n            if AllChem.EmbedMolecule(mol_h, useRandomCoords=True) != 0:\n                continue\n            ids = [0]\n\n        best_conf_score = -1e9\n        for cid in ids:\n            conf = mol_h.GetConformer(cid)\n            lig_pos = np.array([list(conf.GetAtomPosition(a.GetIdx()))\n                                for a in mol_h.GetAtoms()\n                                if a.GetAtomicNum() > 1], dtype=np.float32)\n            if len(lig_pos) == 0:\n                continue\n\n            lig_pos += (pocket_centre - lig_pos.mean(axis=0))\n\n            diff = lig_pos[:, None, :] - rna_arr[None, :, :]\n            dists = np.sqrt((diff ** 2).sum(axis=-1))\n\n            mask_contact = dists < CONTACT_CUTOFF\n            score = np.sum(1.0 / (dists[mask_contact] + 0.5) ** 2)\n\n            n_clashes = np.sum(dists < CLASH_THR)\n            score -= n_clashes * CLASH_PENALTY\n\n            if score > best_conf_score:\n                best_conf_score = score\n\n        total_score += max(best_conf_score, 0.0)\n\n    return float(total_score)\n\ndef codock_rank_samples(samples, seq, smiles_list, plddt_scores=None, delta_g=0.0):\n    if not samples:\n        return samples\n\n    def combined_score(idx):\n        plddt = plddt_scores[idx] if plddt_scores and idx < len(plddt_scores) else 0.5\n        dock  = codock_score_sample(samples[idx], smiles_list) if USE_CODOCK else 0.0\n        return -plddt - 0.002 * dock - 0.001 * abs(delta_g)\n\n    ranked = sorted(range(len(samples)), key=combined_score)\n    return [samples[i] for i in ranked]\n\ndef kabsch_alignment(P, Q):\n    P_centroid = np.mean(P, axis=0)\n    Q_centroid = np.mean(Q, axis=0)\n    H = (P - P_centroid).T @ (Q - Q_centroid)\n    U, S, Vt = np.linalg.svd(H)\n    R = Vt.T @ U.T\n    if np.linalg.det(R) < 0:\n        Vt[2, :] *= -1\n        R = Vt.T @ U.T\n    t = Q_centroid - P_centroid @ R.T\n    return R, t\n\ndef stitch_chunks(chunks, overlap=64):\n    if not chunks: return np.zeros((0, 3))\n    if len(chunks) == 1: return chunks[0]\n    assembled = chunks[0].copy()\n    for i in range(1, len(chunks)):\n        nc = chunks[i].copy()\n        R, t = kabsch_alignment(nc[:overlap], assembled[-overlap:])\n        aligned = (nc @ R.T) + t\n        w = np.linspace(0, 1, overlap)[:, None]\n        assembled[-overlap:] = assembled[-overlap:] * (1 - w) + aligned[:overlap] * w\n        assembled = np.vstack([assembled, aligned[overlap:]])\n    return assembled\n\ndef sort_by_centroid(coords):\n    n = coords.shape[0]\n    if n <= 1: return coords\n    rmsds = np.zeros((n, n))\n    for i in range(n):\n        for j in range(i + 1, n):\n            d = np.sqrt(np.mean(np.sum((coords[i] - coords[j]) ** 2, axis=-1)))\n            rmsds[i, j] = rmsds[j, i] = d\n    return coords[np.argsort(rmsds.sum(1) / (n - 1))]\n\ndef align_assembled_to_vfold(assembled, vfold_ref, min_anchors=10):\n    if assembled.shape != vfold_ref.shape or len(assembled) < min_anchors:\n        return assembled\n    valid = (\n        (np.abs(assembled).sum(axis=1) > 1e-6) &\n        (np.abs(vfold_ref).sum(axis=1)  > 1e-6)\n    )\n    if valid.sum() < min_anchors:\n        return assembled\n    R, t = kabsch_alignment(assembled[valid], vfold_ref[valid])\n    return assembled @ R.T + t\n\n@njit(fastmath=True)\ndef _physics_informed_refinement(coords, n, seg_starts, seg_ends, n_segs, ions_present=True, passes=200, strength=1.0):\n    backbone_force = 0.08 * strength; intrachain_pull = (0.10 if ions_present else 0.02) * strength; interchain_pull = 0.04 * strength\n    atom_seg = np.zeros(n, dtype=np.int64)\n    for s in range(n_segs):\n        for a in range(seg_starts[s], seg_ends[s]): atom_seg[a] = s\n    is_seg_end = np.zeros(n, dtype=np.bool_)\n    for s in range(n_segs):\n        if seg_ends[s] > 0: is_seg_end[seg_ends[s] - 1] = True\n    for _ in range(passes):\n        for i in range(n - 1):\n            if is_seg_end[i]: continue\n            vx = coords[i+1,0]-coords[i,0]; vy = coords[i+1,1]-coords[i,1]; vz = coords[i+1,2]-coords[i,2]\n            d = (vx*vx+vy*vy+vz*vz)**0.5 + 1e-6; f = (5.9 - d) / d * backbone_force\n            coords[i,0] -= vx*f; coords[i,1] -= vy*f; coords[i,2] -= vz*f\n            coords[i+1,0] += vx*f; coords[i+1,1] += vy*f; coords[i+1,2] += vz*f\n        if n > 20:\n            for i in range(n):\n                for j in range(i + 15, n):\n                    dx = coords[i,0]-coords[j,0]; dy = coords[i,1]-coords[j,1]; dz = coords[i,2]-coords[j,2]\n                    dsq = dx*dx + dy*dy + dz*dz\n                    if atom_seg[i] == atom_seg[j]:\n                        if dsq < 64.0:\n                            d = dsq**0.5 + 1e-6; f = (8.0 - d) / d * intrachain_pull\n                            coords[i,0]+=dx*f; coords[i,1]+=dy*f; coords[i,2]+=dz*f\n                            coords[j,0]-=dx*f; coords[j,1]-=dy*f; coords[j,2]-=dz*f\n                    else:\n                        if dsq < 100.0:\n                            d = dsq**0.5 + 1e-6; f = (6.5 - d) / d * interchain_pull\n                            coords[i,0]+=dx*f; coords[i,1]+=dy*f; coords[i,2]+=dz*f\n                            coords[j,0]-=dx*f; coords[j,1]-=dy*f; coords[j,2]-=dz*f\n    return coords\n\ndef adaptive_rna_constraints_with_ions(coords, target_id, segments_map, valid_mask=None, ligand_pull_strength=0.0, confidence=1.0, passes=150, ions_present=True):\n    X = coords.copy(); segments = segments_map.get(target_id, [(0, len(X))])\n    n_atoms = len(X); strength = max(0.75 * (1.0 - min(confidence, 0.97)), 0.05)\n    seg_starts_arr = np.array([s for s, e in segments], dtype=np.int64); seg_ends_arr = np.array([e for s, e in segments], dtype=np.int64)\n    n_segs = len(segments); apply_physics = (confidence < 0.75)\n    for k in range(passes):\n        for s, e in segments:\n            C = X[s:e]; L = e - s\n            if L < 3: continue\n            d = C[1:] - C[:-1]; dist = np.linalg.norm(d, axis=1) + 1e-6\n            adj = d * ((5.95 - dist) / dist)[:, None] * (0.22 * strength)\n            C[:-1] -= adj; C[1:] += adj\n            if L > 2:\n                d2 = C[2:] - C[:-2]; d2n = np.linalg.norm(d2, axis=1) + 1e-6\n                adj2 = d2 * ((10.2 - d2n) / d2n)[:, None] * (0.10 * strength)\n                C[:-2] -= adj2; C[2:] += adj2\n            C[1:-1] += (0.06 * strength) * (0.5 * (C[:-2] + C[2:]) - C[1:-1])\n            if L >= 25:\n                idx = np.linspace(0, L-1, min(L,160)).astype(int) if L > 220 else np.arange(L)\n                P = C[idx]; diff = P[:,None,:] - P[None,:,:]\n                dm = np.linalg.norm(diff, axis=2) + 1e-6; sep = np.abs(idx[:,None] - idx[None,:])\n                mask = (sep > 2) & (dm < 3.2)\n                if np.any(mask):\n                    vec = (diff * ((3.2 - dm) / dm)[:,:,None] * mask[:,:,None]).sum(axis=1)\n                    C[idx] += (0.015 * strength) * vec\n        if apply_physics and k > 60 and (k % 15 == 0):\n            X = _physics_informed_refinement(X, n_atoms, seg_starts_arr, seg_ends_arr, n_segs, ions_present=ions_present, passes=3, strength=strength)\n        if ligand_pull_strength > 0 and valid_mask is not None:\n            gap_mask = ~valid_mask\n            if valid_mask.any() and gap_mask.any():\n                centroid = X[valid_mask].mean(0); diff = centroid - X[gap_mask]\n                dists = np.linalg.norm(diff, axis=1) + 1e-6\n                X[gap_mask] += (diff / dists[:,None]) * (ligand_pull_strength * 0.1 * strength)\n    return X\n\ndef build_configs_local(input_json_path, dump_dir, model_name, cfg):\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(t, p):\n        for k, v in p.items():\n            if isinstance(v, dict) and k in t and isinstance(t[k], dict): deep_update(t[k], v)\n            else: t[k] = v\n\n    deep_update(base, model_configs[model_name])\n    mmcif_dir = os.environ.get(\"PROTENIX_MMCIF_DIR\", \"/kaggle/working/mmcif\")\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 {cfg['USE_MSA']}\",\n        f\"--use_template {cfg['USE_TEMPLATE']}\",\n        f\"--use_rna_msa {cfg['USE_RNA_MSA']}\",\n        f\"--data.template.prot_template_mmcif_dir {mmcif_dir}\",\n        f\"--sample_diffusion.N_sample {cfg['MODEL_N_SAMPLE']}\",\n        f\"--seeds {cfg['SEED']}\",\n        f\"--dtype {cfg['DTYPE']}\",\n        f\"--model.N_cycle {cfg['N_CYCLE']}\",\n        f\"--sample_diffusion.N_step {cfg['N_STEP']}\",\n        f\"--triangle_multiplicative {cfg['TRIMUL_KERNEL']}\",\n        f\"--triangle_attention {cfg['TRIATT_KERNEL']}\",\n    ])\n    return parse_configs(configs=base, arg_str=arg_str, fill_required_with_null=True)\n\n\ndef main():\n    parser = argparse.ArgumentParser()\n    parser.add_argument(\"--gpu-id\",      type=int, required=True)\n    parser.add_argument(\"--json-path\",   required=True)\n    parser.add_argument(\"--queue-pkl\",   required=True)\n    parser.add_argument(\"--segments-pkl\",required=True)\n    parser.add_argument(\"--vfold-pkl\",   required=True)\n    parser.add_argument(\"--result-pkl\",  required=True)\n    parser.add_argument(\"--work-dir\",    required=True)\n    parser.add_argument(\"--cfg-json\",    required=True)\n    args = parser.parse_args()\n\n    with open(args.cfg_json) as f:\n        cfg = json.load(f)\n\n    code_dir = cfg[\"code_dir\"]\n    root_dir = cfg[\"root_dir\"]\n    os.environ[\"PROTENIX_ROOT_DIR\"] = root_dir\n    sys.path.insert(0, code_dir)\n\n    N_SAMPLE      = cfg[\"N_SAMPLE\"]\n    CHUNK_OVERLAP = cfg[\"CHUNK_OVERLAP\"]\n    MODEL_NAME    = cfg[\"MODEL_NAME\"]\n\n    with open(args.queue_pkl, \"rb\") as f:\n        worker_queue = pickle.load(f)\n    with open(args.segments_pkl, \"rb\") as f:\n        segments_map = pickle.load(f)\n    with open(args.vfold_pkl, \"rb\") as f:\n        vfold_arrays = pickle.load(f)\n\n    if not worker_queue:\n        print(f\"[GPU {args.gpu_id}] Empty queue — nothing to do.\")\n        with open(args.result_pkl, \"wb\") as f: pickle.dump({}, f)\n        return\n\n    _dummy = np.random.rand(20, 3).astype(np.float64)\n    _ss = np.array([0], dtype=np.int64); _se = np.array([20], dtype=np.int64)\n    _physics_informed_refinement(_dummy, 20, _ss, _se, 1, True, 2, 0.1)\n\n    from protenix.data.inference.infer_dataloader import InferenceDataset\n    from runner.inference import InferenceRunner, update_gpu_compatible_configs, update_inference_configs\n\n    os.makedirs(args.work_dir, exist_ok=True)\n    configs = build_configs_local(args.json_path, args.work_dir, MODEL_NAME, cfg)\n    configs = update_gpu_compatible_configs(configs)\n\n    gc.collect(); torch.cuda.empty_cache()\n    runner  = InferenceRunner(configs)\n    dataset = InferenceDataset(configs)\n\n    print(f\"[GPU {args.gpu_id}] Starting inference on {len(worker_queue)} target(s), {len(dataset)} chunk(s) total.\")\n\n    raw_chunk_outputs = {tid: {} for tid in worker_queue}\n    raw_chunk_plddts  = {tid: {} for tid in worker_queue}\n    protenix_preds: dict = {}\n\n    for i in tqdm(range(len(dataset)), desc=f\"GPU {args.gpu_id} Diffusion\"):\n        data, atom_array, error_message = dataset[i]\n        chunk_id = data.get(\"sample_name\", f\"sample_{i}\")\n        tid = \"_\".join(chunk_id.split(\"_\")[:-1])\n        try: c_idx = int(chunk_id.split(\"_\")[-1].replace(\"chunk\", \"\"))\n        except Exception: del data, atom_array; gc.collect(); torch.cuda.empty_cache(); continue\n\n        if tid not in worker_queue:\n            del data, atom_array; gc.collect(); torch.cuda.empty_cache(); continue\n\n        info = worker_queue[tid]\n        n_needed      = info[\"n_needed\"]\n        chunk_seq_len = info[\"chunks\"][c_idx][1] - info[\"chunks\"][c_idx][0]\n\n        if error_message:\n            print(f\"  {chunk_id}: data error — {error_message}\")\n            del data, atom_array, error_message; gc.collect(); torch.cuda.empty_cache(); continue\n\n        try:\n            new_cfg = update_inference_configs(configs, data[\"N_token\"].item())\n            new_cfg.sample_diffusion.N_sample = n_needed\n            runner.update_model_configs(new_cfg)\n\n            prediction = runner.predict(data)\n            raw_coords = prediction[\"coordinate\"]\n            feat = data[\"input_feature_dict\"]\n\n            if \"centre_atom_mask\" in feat: mask = (feat[\"centre_atom_mask\"] == 1).to(raw_coords.device)\n            elif \"atom_to_tokatom_idx\" in feat:\n                m11 = (feat[\"atom_to_tokatom_idx\"] == 11).to(raw_coords.device)\n                m12 = (feat[\"atom_to_tokatom_idx\"] == 12).to(raw_coords.device)\n                mask = m11 if abs(m11.sum()-chunk_seq_len) < abs(m12.sum()-chunk_seq_len) else m12\n            else: mask = torch.zeros(raw_coords.shape[1], dtype=torch.bool, device=raw_coords.device)\n\n            coords = raw_coords[:, mask, :].detach().cpu().numpy()\n\n            plddt_scores = None\n            if \"plddt\" in prediction:\n                plddt_raw = prediction[\"plddt\"]\n                plddt_arr = plddt_raw.detach().cpu().numpy() if hasattr(plddt_raw,\"detach\") else np.array(plddt_raw)\n                if plddt_arr.ndim == 2: plddt_scores = plddt_arr.mean(axis=1).tolist()\n                elif plddt_arr.ndim == 1: plddt_scores = [float(plddt_arr.mean())] * coords.shape[0]\n\n            if coords.shape[1] != chunk_seq_len:\n                padded = np.zeros((coords.shape[0], chunk_seq_len, 3), dtype=np.float32)\n                ml = min(coords.shape[1], chunk_seq_len)\n                if ml > 0: padded[:, :ml, :] = coords[:, :ml, :]\n                coords = padded\n\n            raw_chunk_outputs[tid][c_idx] = coords\n            raw_chunk_plddts[tid][c_idx]  = plddt_scores\n\n            # Early assembly\n            n_chunks = len(info[\"chunks\"])\n            if all(ci in raw_chunk_outputs[tid] for ci in range(n_chunks)):\n                _nn = info[\"n_needed\"]; _seq = info[\"seq\"]\n                _assem = []; _pldds = []\n                for s_idx in range(_nn):\n                    sc = [raw_chunk_outputs[tid][ci][s_idx] for ci in range(n_chunks)]\n                    _stitched = stitch_chunks(sc, overlap=CHUNK_OVERLAP)\n                    if tid in vfold_arrays:\n                        _vref = vfold_arrays[tid][info[\"chunks\"][0][0] : info[\"chunks\"][-1][1]]\n                        if len(_vref) == len(_stitched):\n                            _stitched = align_assembled_to_vfold(_stitched, _vref)\n                    _assem.append(_stitched)\n                    cp = raw_chunk_plddts[tid].get(0)\n                    _pldds.append(cp[s_idx] if cp and s_idx < len(cp) else 0.5)\n                \n                # CoDock finalization\n                _row = info[\"row\"]\n                _raw_smis = _row.get(\"ligand_SMILES\", \"\")\n                _raw_smis = \"\" if (not isinstance(_raw_smis, str) or pd.isna(_raw_smis) if not isinstance(_raw_smis, str) else False) else _raw_smis\n                _smis_list = [s.strip() for s in _raw_smis.split(\";\") if s.strip()] if _raw_smis else []\n                _dg = compute_delta_g(_seq)\n                \n                _assem = codock_rank_samples(_assem, _seq, _smis_list, _pldds, _dg)\n                _ac = sort_by_centroid(np.stack(_assem, axis=0))\n                \n                if n_chunks > 1:\n                    _ref_list = []\n                    for s_idx in range(_ac.shape[0]):\n                        _ref_list.append(adaptive_rna_constraints_with_ions(_ac[s_idx], tid, segments_map, confidence=0.95, passes=20, ions_present=False))\n                    _ac = np.stack(_ref_list, axis=0)\n                \n                protenix_preds[tid] = _ac\n                print(f\"  [GPU {args.gpu_id}] {tid}: early-assembled {n_chunks} chunk(s)\")\n                del raw_chunk_outputs[tid], raw_chunk_plddts[tid]; gc.collect()\n\n        except Exception as exc: print(f\"  [GPU {args.gpu_id}] {chunk_id}: Protenix FAILED — {exc}\")\n        finally:\n            if \"feat\" in dir(): del feat\n            if \"prediction\" in dir(): del prediction\n            if \"raw_coords\" in dir(): del raw_coords\n            if \"mask\" in dir(): del mask\n            del data, atom_array; gc.collect(); torch.cuda.empty_cache()\n\n    # Fallback Assembly\n    for tid, info in worker_queue.items():\n        if tid in protenix_preds: continue\n        chunks = info[\"chunks\"]; n_needed = info[\"n_needed\"]; seq = info[\"seq\"]\n        if not all(ci in raw_chunk_outputs.get(tid,{}) for ci in range(len(chunks))):\n            print(f\"  [GPU {args.gpu_id}] {tid}: missing chunks — skipping\"); continue\n        assem = []; pldds = []\n        for si in range(n_needed):\n            sc = [raw_chunk_outputs[tid][ci][si] for ci in range(len(chunks))]\n            stitched = stitch_chunks(sc, overlap=CHUNK_OVERLAP)\n            if tid in vfold_arrays:\n                _vref_k = vfold_arrays[tid][chunks[0][0] : chunks[-1][1]]\n                if len(_vref_k) == len(stitched):\n                    stitched = align_assembled_to_vfold(stitched, _vref_k)\n            assem.append(stitched)\n            cp = raw_chunk_plddts[tid].get(0)\n            pldds.append(cp[si] if cp and si < len(cp) else 0.5)\n            \n        _krow = info[\"row\"]\n        _kraw_smis = _krow.get(\"ligand_SMILES\", \"\")\n        _kraw_smis = \"\" if not isinstance(_kraw_smis, str) else _kraw_smis\n        _ksmis_list = [s.strip() for s in _kraw_smis.split(\";\") if s.strip()] if _kraw_smis else []\n        _kdg = compute_delta_g(seq)\n        \n        assem = codock_rank_samples(assem, seq, _ksmis_list, pldds, _kdg)\n        ac = sort_by_centroid(np.stack(assem, axis=0))\n        \n        if len(chunks) > 1:\n            ref_list = []\n            for s_idx in range(ac.shape[0]):\n                ref_list.append(adaptive_rna_constraints_with_ions(ac[s_idx], tid, segments_map, confidence=0.95, passes=20, ions_present=False))\n            ac = np.stack(ref_list, axis=0)\n            \n        protenix_preds[tid] = ac\n        print(f\"  [GPU {args.gpu_id}] {tid}: Kabsch assembled {len(chunks)} chunk(s)\")\n\n    with open(args.result_pkl, \"wb\") as f: pickle.dump(protenix_preds, f)\n    print(f\"[GPU {args.gpu_id}] Done — saved {len(protenix_preds)} result(s) to {args.result_pkl}\")\n\nif __name__ == \"__main__\":\n    main()\n'''\n    with open(script_path, \"w\", encoding=\"utf-8\") as fh: fh.write(script)\n    print(f\"  GPU worker script written to {script_path}\")\n\n# ─────────────── Main ────────────────────────────────────────────────────────\ndef main() -> None:\n    # JIT warm-up — must use the new segment-array signature\n    _dummy      = np.random.rand(20, 3).astype(np.float64)\n    _seg_starts = np.array([0],  dtype=np.int64)\n    _seg_ends   = np.array([20], dtype=np.int64)\n    _ = apply_ion_stabilization_fast(_dummy, 20)\n    _ = physics_informed_refinement(_dummy, 20, _seg_starts, _seg_ends, 1,\n                                    ions_present=True, passes=2, strength=0.1)\n    \n    test_csv, output_csv, code_dir, root_dir = resolve_paths()\n\n    if not os.path.isdir(code_dir): raise FileNotFoundError(f\"Missing PROTENIX_CODE_DIR: {code_dir}\")\n    os.environ[\"PROTENIX_ROOT_DIR\"] = root_dir\n    sys.path.append(code_dir)\n    ensure_required_files(root_dir)\n    seed_everything(SEED)\n\n    # Wire PDB_RNA CIF files into a writable mmcif dir for Protenix template search\n    mmcif_dir = ensure_mmcif_dir(PDB_RNA_CIF_DIR)\n    os.environ[\"PROTENIX_MMCIF_DIR\"] = mmcif_dir\n\n    test_df_full = pd.read_csv(test_csv)\n    test_df = (test_df_full[test_df_full[\"target_id\"].isin([\"8ZNQ\"])] if not IS_KAGGLE else test_df_full).reset_index(drop=True)\n\n    print(\"\\nLoading training data for TBM …\")\n    train_seqs = pd.read_csv(DEFAULT_TRAIN_CSV)\n    val_seqs = pd.read_csv(DEFAULT_VAL_CSV)\n    train_labels = pd.read_csv(DEFAULT_TRAIN_LBLS)\n    val_labels = pd.read_csv(DEFAULT_VAL_LBLS)\n\n    combined_seqs = pd.concat([train_seqs, val_seqs], ignore_index=True)\n    combined_labels = pd.concat([train_labels, val_labels], ignore_index=True)\n    train_coords = process_labels(combined_labels)\n    segments_map, _ = build_segments_map(test_df)\n\n    # ─── Load rna_metadata.csv for template quality ranking ──────────────────\n    print(\"Loading structural quality metadata...\")\n    quality_map = load_quality_map(METADATA_CSV)\n\n    # ─── NEW: Load PDB SEQRES DB + release dates for expanded template search ─\n    print(\"Loading PDB SEQRES database for expanded template search...\")\n    pdb_seqres_dict = load_pdb_seqres_db(PDB_SEQRES_FASTA)\n    release_dates   = load_pdb_release_dates(PDB_RELEASE_DATES_CSV)\n    # Build k-mer index for fast candidate filtering\n    if pdb_seqres_dict:\n        print(f\"  Building 5-mer index for {len(pdb_seqres_dict):,} SEQRES entries...\")\n        pdb_kmer_index = build_kmer_index(pdb_seqres_dict, k=5)\n        print(f\"  k-mer index ready ({len(pdb_kmer_index):,} unique 5-mers)\")\n    else:\n        pdb_kmer_index = {}\n\n    # ─── Ensure MSA directory is accessible for Protenix ─────────────────────\n    msa_link = \"/kaggle/working/msa\"\n    if os.path.isdir(MSA_DIR) and not os.path.exists(msa_link):\n        try:\n            os.symlink(MSA_DIR, msa_link)\n            print(f\"MSA directory symlinked: {msa_link} → {MSA_DIR}\")\n        except OSError:\n            print(f\"MSA symlink failed (read-only?); Protenix will use default MSA path.\")\n    if os.path.isdir(MSA_DIR):\n        msa_count = len(list(Path(MSA_DIR).glob(\"*.fasta\")))\n        print(f\"MSA files available: {msa_count}\")\n    \n    print(\"Precomputing ligand features for training templates...\")\n    train_ligand_fps = {}\n    for _, row in tqdm(combined_seqs.iterrows(), total=len(combined_seqs), desc=\"Parsing Train Ligands\"):\n        fps, _ = extract_ligand_features(row.get(\"ligand_SMILES\", \"\"), row.get(\"ligand_ids\", \"\"))\n        train_ligand_fps[row[\"target_id\"]] = fps\n\n    # ─── PHASE 1 ─────────────────────────────────────────────────────────────────\n    template_preds, protenix_queue = tbm_phase(test_df, combined_seqs, train_coords, segments_map, train_ligand_fps, quality_map=quality_map, pdb_seqres_dict=pdb_seqres_dict, pdb_kmer_index=pdb_kmer_index, release_dates=release_dates)\n\n    # ─── PHASE 2 ─────────────────────────────────────────────────────────────────\n    protenix_preds: dict = {}   \n\n    if protenix_queue and USE_PROTENIX:\n        print(f\"\\n{'='*60}\\nPHASE 2: Protenix Chunking & Kabsch Assembly for {len(protenix_queue)} targets\\n{'='*60}\")\n\n        work_dir = Path(\"/kaggle/working\")\n        work_dir.mkdir(parents=True, exist_ok=True)\n        \n        # -- Vfold scaffold generation ----------------------------------------\n        # Build global ViennaRNA MFE scaffold for every queued target, write\n        # per-chunk CIFs into mmcif_dir (100% seq-identity -> top template),\n        # and store generators for post-inference global alignment.\n        print(\"Generating Vfold global scaffolds and chunk template CIFs...\")\n        _vfold_generators: dict = {}\n        for _tid, _info in protenix_queue.items():\n            try:\n                _gen = VfoldScaffoldGenerator(_info[\"seq\"], _info[\"strands\"])\n                _gen.generate()\n                _gen.write_cif(\"VF\" + _tid[:2].upper(), mmcif_dir)\n                for _ci, (_cs, _ce) in enumerate(_info[\"chunks\"]):\n                    write_chunk_vfold_cif(_gen, _tid + \"_chunk\" + str(_ci),\n                                          _cs, _ce, mmcif_dir)\n                _vfold_generators[_tid] = _gen\n                print(f\"  {_tid}: Vfold scaffold ready ({len(_info['chunks'])} chunk CIF(s))\")\n            except Exception as _ve:\n                print(f\"  {_tid}: Vfold scaffold FAILED -- {_ve}\")\n        # ---------------------------------------------------------------------\n\n        n_gpus_available = torch.cuda.device_count()\n        n_gpus = min(n_gpus_available, 2)\n        print(f\"  Detected {n_gpus_available} GPU(s) — using {n_gpus} for inference.\")\n\n        worker_cfg = {\n            \"N_SAMPLE\":      N_SAMPLE,\n            \"CHUNK_OVERLAP\": CHUNK_OVERLAP,\n            \"MODEL_NAME\":    MODEL_NAME,\n            \"SEED\":          SEED,\n            \"USE_MSA\":       USE_MSA,\n            \"USE_TEMPLATE\":  USE_TEMPLATE,\n            \"USE_RNA_MSA\":   USE_RNA_MSA,\n            \"MODEL_N_SAMPLE\":MODEL_N_SAMPLE,\n            \"DTYPE\":         DTYPE,\n            \"N_CYCLE\":       N_CYCLE,\n            \"N_STEP\":        N_STEP,\n            \"TRIMUL_KERNEL\": TRIMUL_KERNEL,\n            \"TRIATT_KERNEL\": TRIATT_KERNEL,\n            \"code_dir\":      code_dir,\n            \"root_dir\":      root_dir,\n        }\n        cfg_json_path = str(work_dir / \"worker_cfg.json\")\n        with open(cfg_json_path, \"w\") as _f: json.dump(worker_cfg, _f)\n\n        segments_pkl_path = str(work_dir / \"segments_map.pkl\")\n        with open(segments_pkl_path, \"wb\") as _f: pickle.dump(segments_map, _f)\n        \n        # Save safe NumPy arrays to pass the Vfold data to workers without pickling complex class definitions.\n        vfold_arrays = {tid: gen._coord_array for tid, gen in _vfold_generators.items()}\n\n        if n_gpus >= 2:\n            print(\"  → Multi-GPU mode: Greedy length-based split across workers.\")\n\n            worker_script_path = str(work_dir / \"_protenix_gpu_worker.py\")\n            write_gpu_worker_script(worker_script_path)\n\n            gpu_queues = split_queue_across_gpus(protenix_queue, n_gpus)\n\n            procs = []; result_pkls = []\n            \n            # Serialize Vfold arrays\n            vfold_pkl_path = str(work_dir / \"vfold_arrays.pkl\")\n            with open(vfold_pkl_path, \"wb\") as _f: pickle.dump(vfold_arrays, _f)\n\n            for gpu_id in range(n_gpus):\n                q_subset = gpu_queues[gpu_id]\n                if not q_subset:\n                    print(f\"  GPU {gpu_id}: empty slice — skipping.\")\n                    procs.append(None); result_pkls.append(None)\n                    continue\n\n                gpu_json_path = str(work_dir / f\"input_gpu{gpu_id}.json\")\n                build_input_json_chunked(q_subset, test_df, gpu_json_path)\n\n                gpu_queue_pkl = str(work_dir / f\"queue_gpu{gpu_id}.pkl\")\n                with open(gpu_queue_pkl, \"wb\") as _f: pickle.dump(q_subset, _f)\n\n                result_pkl = str(work_dir / f\"results_gpu{gpu_id}.pkl\")\n                result_pkls.append(result_pkl)\n\n                gpu_work_dir = str(work_dir / f\"outputs_gpu{gpu_id}\")\n                os.makedirs(gpu_work_dir, exist_ok=True)\n\n                cmd = [\n                    sys.executable, worker_script_path,\n                    \"--gpu-id\",       str(gpu_id),\n                    \"--json-path\",    gpu_json_path,\n                    \"--queue-pkl\",    gpu_queue_pkl,\n                    \"--segments-pkl\", segments_pkl_path,\n                    \"--vfold-pkl\",    vfold_pkl_path,\n                    \"--result-pkl\",   result_pkl,\n                    \"--work-dir\",     gpu_work_dir,\n                    \"--cfg-json\",     cfg_json_path,\n                ]\n\n                print(f\"  Launching GPU {gpu_id} worker: {len(q_subset)} target(s) ({sum(len(v['chunks']) for v in q_subset.values())} chunk(s))\")\n\n                env = os.environ.copy()\n                env[\"CUDA_VISIBLE_DEVICES\"] = str(gpu_id)\n                env[\"NUMBA_CACHE_DIR\"] = str(work_dir / f\"numba_cache_gpu{gpu_id}\") \n\n                if \"PROTENIX_MMCIF_DIR\" in os.environ: env[\"PROTENIX_MMCIF_DIR\"] = os.environ[\"PROTENIX_MMCIF_DIR\"]\n\n                proc = subprocess.Popen(cmd, stdout=subprocess.PIPE, stderr=subprocess.STDOUT, text=True, env=env)\n                procs.append(proc)\n\n            try:\n                for gpu_id, proc in enumerate(procs):\n                    if proc is None: continue\n                    stdout, _ = proc.communicate()\n                    print(f\"\\n{'─'*40}\\n  ── GPU {gpu_id} worker output ──\\n{stdout}\")\n                    if proc.returncode != 0: print(f\"  WARNING: GPU {gpu_id} worker exited with code {proc.returncode}\")\n            finally:\n                for proc in procs:\n                    if proc and proc.poll() is None: proc.terminate()\n\n            for gpu_id, result_pkl in enumerate(result_pkls):\n                if result_pkl is None or not os.path.exists(result_pkl):\n                    print(f\"  WARNING: GPU {gpu_id} result file missing.\")\n                    continue\n                with open(result_pkl, \"rb\") as _f: gpu_results = pickle.load(_f)\n                protenix_preds.update(gpu_results)\n                print(f\"  GPU {gpu_id}: merged {len(gpu_results)} result(s)\")\n\n        else:\n            print(\"  → Single-GPU mode: running sequentially on GPU 0.\")\n            # Constraints removed: Protenix V1 does not accept constraint inputs.\n            input_json_path = str(work_dir / \"protenix_chunked_guided.json\")\n            build_input_json_chunked(protenix_queue, test_df, input_json_path)\n\n            from protenix.data.inference.infer_dataloader import InferenceDataset\n            from runner.inference import InferenceRunner, update_gpu_compatible_configs, update_inference_configs\n\n            configs = build_configs(input_json_path, str(work_dir / \"outputs\"), MODEL_NAME)\n            configs = update_gpu_compatible_configs(configs)\n            \n            gc.collect(); torch.cuda.empty_cache()\n            runner = InferenceRunner(configs)\n            dataset = InferenceDataset(configs)\n\n            # Dictionary to store chunked coordinates before assembly\n            raw_chunk_outputs  = {tid: {} for tid in protenix_queue.keys()}\n            # Innovation 4: Store per-sample pLDDT for thermodynamic ranking\n            raw_chunk_plddts   = {tid: {} for tid in protenix_queue.keys()}\n\n            for i in tqdm(range(len(dataset)), desc=\"Protenix Diffusion\"):\n                data, atom_array, error_message = dataset[i]\n                chunk_id = data.get(\"sample_name\", f\"sample_{i}\")\n                \n                # Parse chunk ID\n                tid = \"_\".join(chunk_id.split(\"_\")[:-1])\n                try: c_idx = int(chunk_id.split(\"_\")[-1].replace(\"chunk\", \"\"))\n                except: continue\n\n                if tid not in protenix_queue: continue\n                \n                info = protenix_queue[tid]\n                n_needed = info[\"n_needed\"]\n                chunk_seq_len = info[\"chunks\"][c_idx][1] - info[\"chunks\"][c_idx][0]\n\n                if error_message:\n                    print(f\"  {chunk_id}: data error — {error_message}\")\n                    del data, atom_array, error_message\n                    gc.collect(); torch.cuda.empty_cache()\n                    continue\n\n                try:\n                    new_cfg = update_inference_configs(configs, data[\"N_token\"].item())\n                    new_cfg.sample_diffusion.N_sample = n_needed\n                    runner.update_model_configs(new_cfg)\n\n                    prediction = runner.predict(data)\n                    raw_coords = prediction[\"coordinate\"] \n                    feat = data[\"input_feature_dict\"]\n                    \n                    if \"centre_atom_mask\" in feat: mask = (feat[\"centre_atom_mask\"] == 1).to(raw_coords.device)\n                    elif \"atom_to_tokatom_idx\" in feat:\n                        m11 = (feat[\"atom_to_tokatom_idx\"] == 11).to(raw_coords.device)\n                        m12 = (feat[\"atom_to_tokatom_idx\"] == 12).to(raw_coords.device)\n                        mask = m11 if abs(m11.sum() - chunk_seq_len) < abs(m12.sum() - chunk_seq_len) else m12\n                    else:\n                        mask = torch.zeros(raw_coords.shape[1], dtype=torch.bool, device=raw_coords.device)\n                    \n                    coords = raw_coords[:, mask, :].detach().cpu().numpy()\n                    \n                    # Innovation 4: Extract per-sample pLDDT for thermodynamic ranking\n                    plddt_scores = None\n                    if \"plddt\" in prediction:\n                        plddt_raw = prediction[\"plddt\"]\n                        if hasattr(plddt_raw, 'detach'):\n                            plddt_arr = plddt_raw.detach().cpu().numpy()\n                        else:\n                            plddt_arr = np.array(plddt_raw)\n                        # Average pLDDT per sample across residues\n                        if plddt_arr.ndim == 2:  # (N_sample, N_residue)\n                            plddt_scores = plddt_arr.mean(axis=1).tolist()\n                        elif plddt_arr.ndim == 1:\n                            plddt_scores = [float(plddt_arr.mean())] * coords.shape[0]\n                    \n                    if coords.shape[1] != chunk_seq_len:\n                        padded = np.zeros((coords.shape[0], chunk_seq_len, 3), dtype=np.float32)\n                        min_len = min(coords.shape[1], chunk_seq_len)\n                        if min_len > 0: padded[:, :min_len, :] = coords[:, :min_len, :]\n                        coords = padded\n\n                    raw_chunk_outputs[tid][c_idx] = coords\n                    raw_chunk_plddts[tid][c_idx]  = plddt_scores\n\n                    # ── Early assembly: once ALL chunks for this target are ready,\n                    # assemble immediately and release the per-chunk buffers.\n                    # This prevents raw_chunk_outputs from accumulating every\n                    # target's coords in RAM until the separate Kabsch loop runs.\n                    n_chunks = len(info[\"chunks\"])\n                    if all(ci in raw_chunk_outputs[tid] for ci in range(n_chunks)):\n                        _info   = protenix_queue[tid]\n                        _seq    = _info[\"seq\"]\n                        _segs   = _info[\"segments\"]\n                        _nchk   = _info[\"chunks\"]\n                        _nn     = _info[\"n_needed\"]\n                        _assem  = []\n                        _pldds  = []\n                        for s_idx in range(_nn):\n                            _sample_chunks = [raw_chunk_outputs[tid][ci][s_idx]\n                                              for ci in range(n_chunks)]\n                            _stitched = stitch_chunks(_sample_chunks, overlap=CHUNK_OVERLAP)\n                            # Global Vfold alignment -- corrects multi-chunk drift\n                            if _tid in _vfold_generators:\n                                _vref = _vfold_generators[_tid].chunk_coords(\n                                    info[\"chunks\"][0][0], info[\"chunks\"][-1][1])\n                                if len(_vref) == len(_stitched):\n                                    _stitched = align_assembled_to_vfold(_stitched, _vref)\n                            _assem.append(_stitched)\n                            _cp = raw_chunk_plddts[tid].get(0)\n                            _pldds.append(_cp[s_idx] if _cp and s_idx < len(_cp) else 0.5)\n\n                        # CoDock finalization: re-rank using pLDDT + ΔG + ligand pocket score\n                        _row      = protenix_queue[tid][\"row\"]\n                        # Use pd.isna guard — row values may be NaN floats, not empty strings\n                        _raw_smis = _row.get(\"ligand_SMILES\", \"\")\n                        _raw_smis = \"\" if (not isinstance(_raw_smis, str) or pd.isna(_raw_smis) if not isinstance(_raw_smis, str) else False) else _raw_smis\n                        _smis_list = [s.strip() for s in _raw_smis.split(\";\") if s.strip()] if _raw_smis else []\n                        _dg        = compute_delta_g(_seq)\n                        _assem     = codock_rank_samples(_assem, _seq, _smis_list, _pldds, _dg)\n                        _ac    = np.stack(_assem, axis=0)\n                        _ac    = sort_by_centroid(_ac)\n\n                        if n_chunks > 1:\n                            _ref = []\n                            for s_idx in range(_ac.shape[0]):\n                                _ref.append(adaptive_rna_constraints_with_ions(\n                                    _ac[s_idx], tid, segments_map,\n                                    confidence=0.95, passes=20, ions_present=False,\n                                ))\n                            _ac = np.stack(_ref, axis=0)\n\n                        protenix_preds[tid] = _ac\n                        print(f\"  {tid}: early-assembled {n_chunks} chunk(s) → freed chunk buffers\")\n                        # Free chunk buffers immediately — don't hold until Kabsch loop\n                        del raw_chunk_outputs[tid], raw_chunk_plddts[tid]\n                        gc.collect()\n\n                except Exception as exc:\n                    print(f\"  {chunk_id}: Protenix FAILED — {exc}\")\n\n                finally:\n                    # feat holds a reference to input_feature_dict keeping GPU\n                    # tensors alive until the next GC sweep — free it explicitly.\n                    if 'feat' in dir(): del feat\n                    if 'prediction' in dir(): del prediction\n                    if 'raw_coords' in dir(): del raw_coords\n                    if 'mask' in dir(): del mask\n                    del data, atom_array\n                    gc.collect(); torch.cuda.empty_cache()\n\n            # ─── Kabsch Assembly, Thermodynamic Ranking & Consensus ───\n            # Targets assembled early (all chunks resolved inside the loop above)\n            # are already in protenix_preds and their buffers are freed.\n            # This loop only handles targets that had chunk failures mid-way.\n            print(\"\\nExecuting SVD 3D Alignment, ΔG Thermodynamic Ranking & Consensus Sort...\")\n            for tid, info in protenix_queue.items():\n                # Skip: already assembled in the early-assembly path above.\n                if tid in protenix_preds:\n                    continue\n                n_needed = info[\"n_needed\"]\n                chunks = info[\"chunks\"]\n                seq = info[\"seq\"]\n                \n                # Ensure all chunks folded successfully\n                if not all(c_idx in raw_chunk_outputs[tid] for c_idx in range(len(chunks))):\n                    print(f\"  {tid}: Missing chunk folds. Falling back.\")\n                    protenix_preds[tid] = None\n                    continue\n                    \n                assembled_samples = []\n                assembled_plddts  = []\n                for sample_idx in range(n_needed):\n                    sample_chunks = [raw_chunk_outputs[tid][c_idx][sample_idx] for c_idx in range(len(chunks))]\n                    stitched = stitch_chunks(sample_chunks, overlap=CHUNK_OVERLAP)\n                    # Global Vfold alignment -- corrects multi-chunk drift\n                    if tid in _vfold_generators:\n                        _vref_k = _vfold_generators[tid].chunk_coords(\n                            chunks[0][0], chunks[-1][1])\n                        if len(_vref_k) == len(stitched):\n                            stitched = align_assembled_to_vfold(stitched, _vref_k)\n                    assembled_samples.append(stitched)\n                    \n                    # Aggregate pLDDT across chunks (mean of first chunk's pLDDT as proxy)\n                    chunk_plddt = raw_chunk_plddts[tid].get(0)\n                    if chunk_plddt and sample_idx < len(chunk_plddt):\n                        assembled_plddts.append(chunk_plddt[sample_idx])\n                    else:\n                        assembled_plddts.append(0.5)\n\n                # CoDock finalization: re-rank using pLDDT + ΔG + ligand pocket score\n                _krow      = test_df[test_df[\"target_id\"] == tid].iloc[0]\n                _kraw_smis = _krow.get(\"ligand_SMILES\", \"\")\n                _kraw_smis = \"\" if not isinstance(_kraw_smis, str) else _kraw_smis\n                _ksmis_list = [s.strip() for s in _kraw_smis.split(\";\") if s.strip()] if _kraw_smis else []\n                _kdg        = compute_delta_g(seq)\n                assembled_samples = codock_rank_samples(\n                    assembled_samples, seq, _ksmis_list, assembled_plddts, _kdg\n                )\n                assembled_coords = np.stack(assembled_samples, axis=0)  # (N, L, 3)\n\n                # Geometric consensus sort on top of thermodynamic ranking\n                assembled_coords = sort_by_centroid(assembled_coords)\n\n                # Do NOT apply adaptive_rna_constraints_with_ions to Protenix predictions.\n                # Only apply a very light boundary-heal pass for multi-chunk assemblies.\n                if len(chunks) > 1:\n                    refined_coords = []\n                    for s_idx in range(assembled_coords.shape[0]):\n                        ref_c = adaptive_rna_constraints_with_ions(\n                            assembled_coords[s_idx], tid, segments_map,\n                            confidence=0.95, passes=20, ions_present=False,\n                        )\n                        refined_coords.append(ref_c)\n                    assembled_coords = np.stack(refined_coords, axis=0)\n\n                protenix_preds[tid] = assembled_coords\n                print(f\"  {tid}: assembled {len(chunks)} chunk(s), {len(info['strands'])} chain type(s) → Kabsch + ΔG ranking ✓\")\n\n    # ─── PHASE 3 ─────────────────────────────────────────────────────────────────\n    print(f\"\\n{'='*60}\\nPHASE 3: Combine Hybrid TBM + Protenix + De-novo\\n{'='*60}\")\n    all_rows = []\n\n    for _, row in test_df.iterrows():\n        tid, seq = row[\"target_id\"], row[\"sequence\"]\n        combined = list(template_preds.get(tid, [])) \n\n        ptx = protenix_preds.get(tid)\n        if ptx is not None and ptx.ndim == 3:\n            for j in range(ptx.shape[0]):\n                if len(combined) >= N_SAMPLE: break\n                combined.append(ptx[j])\n\n        while len(combined) < N_SAMPLE:\n            seed_val = row.name * 1000000 + len(combined) * 1000\n            # Innovation 2: Use torsion-space helix for de-novo fallback\n            # Vfold scaffold embeds ViennaRNA 2D priors -- far better\n            # de-novo starting geometry than a random helix.\n            if USE_PROTENIX and tid in _vfold_generators:\n                dn = _vfold_generators[tid].chunk_coords(0, len(seq))\n            else:\n                dn = generate_rna_structure(seq, seed=seed_val)\n            combined.append(adaptive_rna_constraints_with_ions(dn, tid, segments_map, confidence=0.2, passes=150))\n\n        stacked = np.stack(combined[:N_SAMPLE], axis=0)\n        all_rows.extend(coords_to_rows(tid, seq, stacked))\n\n    sub = pd.DataFrame(all_rows)\n    cols = [\"ID\", \"resname\", \"resid\"] + [f\"{c}_{i}\" for i in range(1, N_SAMPLE + 1) for c in [\"x\", \"y\", \"z\"]]\n    coord_cols = [c for c in cols if c.startswith((\"x_\", \"y_\", \"z_\"))]\n    sub[coord_cols] = sub[coord_cols].clip(-999.999, 9999.999)\n    sub[cols].to_csv(output_csv, index=False)\n    print(f\"\\n✓ Saved submission to {output_csv}  ({len(sub):,} rows)\")\n\nif __name__ == \"__main__\":\n    main()","metadata":{"execution":{"iopub.status.busy":"2026-03-14T03:57:07.012305Z","iopub.execute_input":"2026-03-14T03:57:07.012587Z","iopub.status.idle":"2026-03-14T04:00:33.515762Z","shell.execute_reply.started":"2026-03-14T03:57:07.012561Z","shell.execute_reply":"2026-03-14T04:00:33.514951Z"},"papermill":{"duration":254.776681,"end_time":"2026-03-12T21:18:14.251395","exception":false,"start_time":"2026-03-12T21:13:59.474714","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==============================================================================\n# LOCAL TESTING & VISUALIZATION CELL \n# (Will automatically skip during Kaggle submission runs)\n# ==============================================================================\n\nimport os\nimport re\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport subprocess\nimport shutil\n\nif not IS_KAGGLE:\n    print(\"\\n\" + \"=\"*60)\n    print(\"STARTING LOCAL EVALUATION & VISUALIZATION\")\n    print(\"=\"*60)\n\n    # --- Config ---\n    USALIGN_BIN = \"/kaggle/working/USalign_bin\"\n\n    if not os.path.exists(USALIGN_BIN):\n        os.system(\"cp /kaggle/input/datasets/metric/usalign/USalign /kaggle/working/USalign_bin\")\n        os.system(\"chmod +x /kaggle/working/USalign_bin\")\n\n    EVAL_DIR = \"/kaggle/working/eval_tmp\"\n    os.makedirs(EVAL_DIR, exist_ok=True)\n\n    # --- PDB Writing Helpers ---\n    def sanitize(xyz): return min(max(xyz,-999.999),9999.999)\n\n    def write_target_line(atom_name, atom_serial, residue_name, chain_id, residue_num,\n                          x_coord, y_coord, z_coord, occupancy=1.0, b_factor=0.0, atom_type='P') -> str:\n        return f'ATOM  {atom_serial:>5d}  {atom_name:4s}{residue_name:>3s} {chain_id:1s}{residue_num:>4d}    {sanitize(x_coord):>8.3f}{sanitize(y_coord):>8.3f}{sanitize(z_coord):>8.3f}{occupancy:>6.2f}{b_factor:>6.2f}           {atom_type}\\n'\n\n    def write_df_to_pdb(df, suffix, out_path, is_native=False):\n        # Sort by residue ID to maintain proper sequential connections\n        df_sorted = df.copy()\n        df_sorted['__resid_int'] = df_sorted['resid'].astype(int)\n        df_sorted = df_sorted.sort_values('__resid_int').reset_index(drop=True)\n        \n        resolved_cnt = 0\n        with open(out_path, 'w') as fh:\n            for _, row in df_sorted.iterrows():\n                x = row[f'x_{suffix}']\n                y = row[f'y_{suffix}']\n                z = row[f'z_{suffix}']\n                \n                # Skip invalid/unresolved coordinates\n                if x > -1e5 and y > -1e5 and z > -1e5 and not np.isnan(x):\n                    resolved_cnt += 1\n                    resid_num = int(row['resid'])\n                    resname = row['resname'] if 'resname' in row else 'A'\n                    fh.write(write_target_line(\"C1'\", resid_num, resname, 'A', resid_num, x, y, z, atom_type='C'))\n        return resolved_cnt\n\n    def parse_tmscore(output: str) -> float:\n        matches = re.findall(r'TM-score=\\s+([\\d.]+)', output)\n        return float(matches[1]) if len(matches) > 1 else 0.0\n\n    # --- Main Evaluation Logic ---\n    print(\"Loading predictions and ground truth...\")\n    sub_df = pd.read_csv(DEFAULT_OUTPUT)\n    val_lbl_df = pd.read_csv(DEFAULT_VAL_LBLS)\n\n    # Load validation sequences to map sequence lengths for output visibility\n    val_seq_df = pd.read_csv(DEFAULT_VAL_CSV)\n    val_seq_df['seq_len'] = val_seq_df['sequence'].apply(len)\n    len_map = val_seq_df.set_index('target_id')['seq_len'].to_dict()\n\n    # Extract unique target IDs from the submission ID column\n    sub_df['target_id'] = sub_df['ID'].str.split('_').str[0]\n    targets = sub_df['target_id'].unique()\n\n    for tid in targets:\n        seq_len = len_map.get(tid, \"Unknown\")\n        print(f\"\\nEvaluating Target: {tid} (Length: {seq_len} nt)\")\n        \n        pred_target_df = sub_df[sub_df['target_id'] == tid].copy()\n        native_target_df = val_lbl_df[val_lbl_df['ID'].str.startswith(f\"{tid}_\")].copy()\n        \n        if native_target_df.empty:\n            print(f\"  -> Skipping {tid}: No native ground truth found in validation labels.\")\n            continue\n\n        native_pdb = os.path.join(EVAL_DIR, f\"{tid}_native.pdb\")\n        native_atoms = write_df_to_pdb(native_target_df, \"1\", native_pdb, is_native=True)\n        \n        tm_scores = []\n        plot_data = [] \n        \n        for i in range(1, 6):\n            pred_pdb = os.path.join(EVAL_DIR, f\"{tid}_pred_{i}.pdb\")\n            pred_atoms = write_df_to_pdb(pred_target_df, str(i), pred_pdb)\n            \n            if pred_atoms > 2 and native_atoms > 2:\n                cmd = f'{USALIGN_BIN} {pred_pdb} {native_pdb} -atom \" C1\\'\" -TMscore 1'\n                result = os.popen(cmd).read()\n                score = parse_tmscore(result)\n            else:\n                score = 0.0\n                \n            tm_scores.append(score)\n            \n            valid_coords = pred_target_df[(pred_target_df[f'x_{i}'] > -1000) & (pred_target_df[f'x_{i}'].notna())]\n            coords = valid_coords[[f'x_{i}', f'y_{i}', f'z_{i}']].values\n            plot_data.append((coords, f\"Sample {i}\\nTM-Score: {score:.4f}\"))\n\n        print(f\"  -> TM-Scores: {[round(s, 4) for s in tm_scores]}\")\n        \n        valid_native = native_target_df[(native_target_df['x_1'] > -1000) & (native_target_df['x_1'].notna())]\n        native_coords = valid_native[['x_1', 'y_1', 'z_1']].values\n        plot_data.append((native_coords, \"Ground Truth\\n(Native)\"))\n\n        # --- 3D Visualization ---\n        fig = plt.figure(figsize=(20, 10))\n        fig.suptitle(f\"Target: {tid} ({seq_len} nt) | Best TM-Score: {max(tm_scores):.4f}\", fontsize=16)\n        \n        for idx, (coords, title) in enumerate(plot_data):\n            ax = fig.add_subplot(2, 3, idx + 1, projection='3d')\n            \n            if len(coords) > 0:\n                c_len = len(coords)\n                colors = np.arange(c_len)\n                \n                ax.plot(coords[:, 0], coords[:, 1], coords[:, 2], color='gray', alpha=0.5, linewidth=1)\n                sc = ax.scatter(coords[:, 0], coords[:, 1], coords[:, 2], \n                                c=colors, cmap='viridis', s=20, alpha=0.8)\n                \n                ax.set_xticklabels([])\n                ax.set_yticklabels([])\n                ax.set_zticklabels([])\n            \n            ax.set_title(title)\n            \n        plt.tight_layout()\n        plt.show()\n\n    # Cleanup temporary files\n    shutil.rmtree(EVAL_DIR, ignore_errors=True)","metadata":{"execution":{"iopub.status.busy":"2026-03-14T04:00:33.517471Z","iopub.execute_input":"2026-03-14T04:00:33.517753Z","iopub.status.idle":"2026-03-14T04:00:35.213579Z","shell.execute_reply.started":"2026-03-14T04:00:33.51773Z","shell.execute_reply":"2026-03-14T04:00:35.212902Z"},"papermill":{"duration":2.193754,"end_time":"2026-03-12T21:18:16.450859","exception":false,"start_time":"2026-03-12T21:18:14.257105","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null}]}