{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":14787388,"datasetId":9453383,"databundleVersionId":15641141},{"sourceType":"datasetVersion","sourceId":14786962,"datasetId":9447634,"databundleVersionId":15640661},{"sourceType":"datasetVersion","sourceId":11118830,"datasetId":6933267,"databundleVersionId":11511771},{"sourceType":"datasetVersion","sourceId":14874339,"datasetId":9502242,"databundleVersionId":15736806},{"sourceType":"datasetVersion","sourceId":14519720,"datasetId":9271415,"databundleVersionId":15347344},{"sourceType":"datasetVersion","sourceId":10855324,"datasetId":6742586,"databundleVersionId":11219268},{"sourceType":"modelInstanceVersion","sourceId":781970,"databundleVersionId":16019818,"modelInstanceId":596642,"modelId":608905},{"sourceType":"modelInstanceVersion","sourceId":311741,"databundleVersionId":11641144,"modelInstanceId":264400,"modelId":285488},{"sourceType":"kernelVersion","sourceId":304072227}],"dockerImageVersionId":31260,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"! python -m pip install -q --no-index --find-links=/kaggle/input/notebooks/honganzhu/biopython-whl -r /kaggle/input/notebooks/honganzhu/biopython-whl/requirements.txt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T17:18:01.238498Z","iopub.execute_input":"2026-03-09T17:18:01.239204Z","iopub.status.idle":"2026-03-09T17:18:12.374368Z","shell.execute_reply.started":"2026-03-09T17:18:01.239164Z","shell.execute_reply":"2026-03-09T17:18:12.373346Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ═══════════════════════════════════════════════════════════════════════════════\n# §0  Environment & Installation\n# ═══════════════════════════════════════════════════════════════════════════════\n\nimport os\n\nos.environ[\"PYTHONHASHSEED\"] = \"42\"\nos.environ[\"CUBLAS_WORKSPACE_CONFIG\"] = \":4096:8\"\n\n# ── Install RNAPro ────────────────────────────────────────────────────────────\nos.system(\"cp -r /kaggle/input/datasets/theoviel/rnapro-src/RNAPro /kaggle/working/RNAPro\")\nos.system(\"cp /kaggle/input/datasets/theoviel/rnapro-src/rnapro-private-best-500m.ckpt /kaggle/working/\")\nos.chdir(\"/kaggle/working/RNAPro\")\nos.system(\"pip install -e . --no-deps\")\nos.chdir(\"/kaggle/working\")\n\nRNAPRO_CCD_DIR = \"/kaggle/working/RNAPro/release_data/ccd_cache/\"\nos.makedirs(RNAPRO_CCD_DIR, exist_ok=True)\nos.system(f\"cp /kaggle/input/datasets/geraseva/protenix-checkpoints/components.v20240608.cif {RNAPRO_CCD_DIR}\")\nos.system(f\"cp /kaggle/input/datasets/geraseva/protenix-checkpoints/components.v20240608.cif.rdkit_mol.pkl {RNAPRO_CCD_DIR}\")\nprint(\"✓ RNAPro installed\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# BLOCK 1: NEW CELL — Add AFTER §0 (after RNAPro install, before §1 Imports)\n# =============================================================================\n\n# ── Install Boltz2 ────────────────────────────────────────────────────────────\nos.system(\"cp -r /kaggle/input/datasets/lbugnon/boltz-src-minimal /kaggle/working/boltz-src-minimal\")\nos.system(\"pip install --no-index --no-build-isolation -e /kaggle/working/boltz-src-minimal\")\nos.system(\"mkdir -p /kaggle/working/boltz_cache\")\nos.system(\"cp -r /kaggle/input/datasets/lbugnon/boltz2 /kaggle/working/boltz_cache\")\nos.system(\"bash -c 'if [ -d /kaggle/working/boltz_cache/boltz2/mols/mols ]; then find /kaggle/working/boltz_cache/boltz2/mols/mols -maxdepth 1 -type f -print0 | xargs -0 mv -t /kaggle/working/boltz_cache/boltz2/mols/ && rm -rf /kaggle/working/boltz_cache/boltz2/mols/mols/; fi'\")\nos.system(\"tar -cf /kaggle/working/boltz_cache/boltz2/mols.tar -C /kaggle/working/boltz_cache/boltz2 mols\")\nos.system(\"pip install biopandas --quiet\")\nprint(\"✓ Boltz2 installed\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ═══════════════════════════════════════════════════════════════════════════════\n# §1  Imports   (FIXED — added shutil)\n# ═══════════════════════════════════════════════════════════════════════════════\n\nimport gc\nimport hashlib\nimport json\nimport shutil\nimport sys\nimport time\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom Bio.Align import PairwiseAligner\nfrom tqdm import tqdm\n\nos.environ[\"LAYERNORM_TYPE\"] = \"torch\"\nos.environ.setdefault(\"RNA_MSA_DEPTH_LIMIT\", \"512\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ═══════════════════════════════════════════════════════════════════════════════\n# §2  Constants & Configuration   (CHANGED — Boltz2 is for ensembling, not slots)\n# ═══════════════════════════════════════════════════════════════════════════════\n\nIS_KAGGLE = bool(os.environ.get(\"KAGGLE_IS_COMPETITION_RERUN\", \"\"))\nLOCAL_N_SAMPLES = 2\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 the first {LOCAL_N_SAMPLES} test targets \"\n          f\"will be processed to save time.\")\n\nDATA_BASE          = \"/kaggle/input/stanford-rna-3d-folding-2\"\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\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          = 719\nMAX_SEQ_LEN   = int(os.environ.get(\"MAX_SEQ_LEN\", \"512\"))\nCHUNK_OVERLAP = int(os.environ.get(\"CHUNK_OVERLAP\", \"64\"))\n\nMIN_SIMILARITY       = float(os.environ.get(\"MIN_SIMILARITY\", \"0.0\"))\nMIN_PERCENT_IDENTITY = float(os.environ.get(\"MIN_PERCENT_IDENTITY\", \"0.0\"))\n\nCHAIN_TRUST_MIN_PCT = 0.0\nFULL_LENGTH_RATIO   = 0.3\nCHAIN_LENGTH_RATIO  = 0.5\n\nUSE_PROTENIX = True\nUSE_RNAPRO   = True\n\nMIN_PROTENIX_SLOTS = 1                          # at least 1 model-predicted slot\nMAX_TBM_SLOTS      = N_SAMPLE - MIN_PROTENIX_SLOTS   # = 4\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\"}:\n        return \"true\"\n    if v in {\"0\", \"false\", \"f\", \"no\", \"n\", \"off\"}:\n        return \"false\"\n    return \"true\" if default else \"false\"\n\nUSE_MSA      = _parse_bool(os.environ.get(\"USE_MSA\", \"true\"))\nUSE_TEMPLATE = _parse_bool(os.environ.get(\"USE_TEMPLATE\", \"false\"))\nUSE_RNA_MSA  = _parse_bool(os.environ.get(\"USE_RNA_MSA\", \"true\"))\n\nMODEL_N_SAMPLE = int(os.environ.get(\"MODEL_N_SAMPLE\", str(N_SAMPLE)))\n\nRNAPRO_MAX_LEN       = 1000\nRNAPRO_CHECKPOINT    = \"/kaggle/working/rnapro-private-best-500m.ckpt\"\nRNAPRO_DUMP_DIR      = \"/kaggle/working/rnapro_output\"\nRNAPRO_MODEL_NAME    = \"rnapro_base\"\nRNAPRO_SEQUENCES_CSV = \"/kaggle/working/rnapro_sequences.csv\"\n\nQUALITY_WEIGHT    = float(os.environ.get(\"QUALITY_WEIGHT\", \"0.4\"))\n\nos.environ[\"N_GPUS\"] = \"2\"\nN_GPUS = int(os.environ.get(\"N_GPUS\", \"0\"))\n\ndef _get_n_gpus() -> int:\n    if N_GPUS > 0:\n        return min(N_GPUS, torch.cuda.device_count())\n    return max(torch.cuda.device_count(), 1)\n\n# ── Boltz2 — used for ensembling with Protenix, NOT as separate slot ──────────\nUSE_BOLTZ2        = True\nBOLTZ_CACHE_DIR   = \"/kaggle/working/boltz_cache/boltz2/\"\nBOLTZ_MAX_TOKENS  = 900\nBOLTZ_INPUT_DIR   = \"/kaggle/working/boltz_input\"\nBOLTZ_OUTPUT_BASE = \"/kaggle/working/boltz_repeat\"\n\n# ═══════════════════════════════════════════════════════════════════════════════\n# §2 addition — Add AFTER the existing §2 cell (after BOLTZ_ENSEMBLE_WEIGHT)\n# ═══════════════════════════════════════════════════════════════════════════════\n\ndef get_slot_allocation(seq_len: int) -> tuple:\n    \"\"\"\n    Returns (max_tbm_slots, min_protenix_slots) based on sequence length.\n\n    - Short (≤ MAX_SEQ_LEN):  max 2 TBM, at least 3 Protenix\n    - Long  (> MAX_SEQ_LEN):  up to 5 TBM, 0 minimum Protenix (TBM-first)\n    \"\"\"\n    if seq_len > MAX_SEQ_LEN:\n        return N_SAMPLE-1, 1      # up to 5 TBM, 0 minimum Protenix\n    else:\n        return 2, 3             # max 2 TBM, at least 3 Protenix","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"CUDA_VISIBLE_DEVICES:\", os.environ.get(\"CUDA_VISIBLE_DEVICES\"))\nprint(\"is_available:\", torch.cuda.is_available())\nprint(\"device_count:\", torch.cuda.device_count())\nfor i in range(torch.cuda.device_count()):\n    print(i, torch.cuda.get_device_name(i))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ═══════════════════════════════════════════════════════════════════════════════\n# §3  General Utilities\n# ═══════════════════════════════════════════════════════════════════════════════\n\ndef stable_hash(s: str) -> int:\n    \"\"\"Deterministic hash that doesn't change between Python sessions.\"\"\"\n    return int(hashlib.sha256(s.encode()).hexdigest(), 16) % (2**63)\n\n\ndef seed_everything(seed: int) -> None:\n    import random\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    os.environ[\"CUBLAS_WORKSPACE_CONFIG\"] = \":4096:8\"\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(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\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\n\ndef ensure_required_files(root_dir: str) -> None:\n    checks = [\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    for p, name in checks:\n        if not p.exists():\n            raise FileNotFoundError(f\"Missing {name}: {p}\")\n\n\ndef pad_samples(coords: np.ndarray, n: int) -> np.ndarray:\n    if coords.shape[0] >= n:\n        return coords[:n]\n    if coords.shape[0] == 0:\n        return np.zeros((n, coords.shape[1], 3), dtype=coords.dtype)\n    extra = np.repeat(coords[:1], n - coords.shape[0], axis=0)\n    return np.concatenate([coords, extra], axis=0)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ═══════════════════════════════════════════════════════════════════════════════\n# §4  Protenix Input / Config Helpers\n# ═══════════════════════════════════════════════════════════════════════════════\n\ndef build_input_json(df: pd.DataFrame, json_path: str) -> None:\n    data = [\n        {\n            \"name\": row[\"target_id\"],\n            \"covalent_bonds\": [],\n            \"sequences\": [{\"rnaSequence\": {\"sequence\": row[\"sequence\"], \"count\": 1}}],\n        }\n        for _, row in df.iterrows()\n    ]\n    with open(json_path, \"w\", encoding=\"utf-8\") as f:\n        json.dump(data, f)\n\n\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(target, patch):\n        for k, v in patch.items():\n            if isinstance(v, dict) and k in target and isinstance(target[k], dict):\n                _deep_update(target[k], v)\n            else:\n                target[k] = v\n\n    _deep_update(base, model_configs[model_name])\n    arg_str = \" \".join([\n        f\"--model_name {model_name}\",\n        f\"--input_json_path {input_json_path}\",\n        f\"--dump_dir {dump_dir}\",\n        f\"--use_msa {USE_MSA}\",\n        f\"--use_template {USE_TEMPLATE}\",\n        f\"--use_rna_msa {USE_RNA_MSA}\",\n        f\"--sample_diffusion.N_sample {MODEL_N_SAMPLE}\",\n        f\"--seeds {SEED}\",\n    ])\n    return parse_configs(configs=base, arg_str=arg_str, fill_required_with_null=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ═══════════════════════════════════════════════════════════════════════════════\n# §5  C1' Mask Extraction\n# ═══════════════════════════════════════════════════════════════════════════════\n\ndef get_c1_mask(data: dict, atom_array) -> torch.Tensor:\n    \"\"\"Extract C1' atom mask from atom_array or fallback to feature dict.\"\"\"\n    # 1. Try atom_array attributes\n    if atom_array is not None:\n        try:\n            if hasattr(atom_array, \"centre_atom_mask\"):\n                m = atom_array.centre_atom_mask == 1\n                if hasattr(atom_array, \"is_rna\"):\n                    m = m & atom_array.is_rna\n                return torch.from_numpy(m).bool()\n            if hasattr(atom_array, \"atom_name\"):\n                base = atom_array.atom_name == \"C1'\"\n                if hasattr(atom_array, \"is_rna\"):\n                    base = base & atom_array.is_rna\n                return torch.from_numpy(base).bool()\n        except Exception:\n            pass\n\n    # 2. Fallback to feature dict\n    f = data[\"input_feature_dict\"]\n    if \"center_atom_mask\" in f:\n        return (f[\"center_atom_mask\"] == 1).bool()\n    if \"centre_atom_mask\" in f:\n        return (f[\"centre_atom_mask\"] == 1).bool()\n\n    # 3. Heuristic: pick atom_to_tokatom_idx closest to N_token\n    n_tokens = data.get(\"N_token\", torch.tensor(0)).item()\n    m11 = (f[\"atom_to_tokatom_idx\"] == 11).bool()\n    m12 = (f[\"atom_to_tokatom_idx\"] == 12).bool()\n    if abs(m11.sum().item() - n_tokens) < abs(m12.sum().item() - n_tokens):\n        return m11\n    return m12\n\n\ndef get_feature_c1_mask(data: dict) -> torch.Tensor:\n    f = data[\"input_feature_dict\"]\n    if \"centre_atom_mask\" in f:\n        return f[\"centre_atom_mask\"].long() == 1\n    return f[\"atom_to_tokatom_idx\"].long() == 12","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ═══════════════════════════════════════════════════════════════════════════════\n# §6  Submission Formatting\n# ═══════════════════════════════════════════════════════════════════════════════\n\ndef coords_to_rows(target_id: str, seq: str, coords: np.ndarray) -> list:\n    \"\"\"Convert coords (N_SAMPLE, seq_len, 3) to submission rows.\"\"\"\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]:\n                x, y, z = coords[s, i]\n            else:\n                x, y, z = 0.0, 0.0, 0.0\n            row[f\"x_{s + 1}\"] = float(x)\n            row[f\"y_{s + 1}\"] = float(y)\n            row[f\"z_{s + 1}\"] = float(z)\n        rows.append(row)\n    return rows\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ═══════════════════════════════════════════════════════════════════════════════\n# §7  Sequence Alignment & TBM Helpers  (FULL REPLACEMENT)\n# ═══════════════════════════════════════════════════════════════════════════════\n\nimport re\n\ndef _make_aligner() -> PairwiseAligner:\n    al = PairwiseAligner()\n    al.mode                          = \"global\"\n    al.match_score                   = 2\n    al.mismatch_score                = -1.5\n    al.open_gap_score                = -8\n    al.extend_gap_score              = -0.4\n    al.query_left_open_gap_score     = -8\n    al.query_left_extend_gap_score   = -0.4\n    al.query_right_open_gap_score    = -8\n    al.query_right_extend_gap_score  = -0.4\n    al.target_left_open_gap_score    = -8\n    al.target_left_extend_gap_score  = -0.4\n    al.target_right_open_gap_score   = -8\n    al.target_right_extend_gap_score = -0.4\n    return al\n\n_aligner = _make_aligner()\n\n\ndef query_chain_key(qc):\n    \"\"\"Unique key for a query-chain instance.\"\"\"\n    return (qc[\"chain_id\"], qc[\"copy\"], qc[\"start\"], qc[\"end\"])\n\n\ndef parse_fasta_with_auth(all_sequences_str):\n    sequences, auth_to_std = {}, {}\n    if pd.isna(all_sequences_str):\n        return sequences, auth_to_std\n    lines = all_sequences_str.strip().split('\\n')\n    current_chains, current_seq = [], []\n    for line in lines:\n        line = line.strip()\n        if line.startswith('>'):\n            if current_chains and current_seq:\n                seq = ''.join(current_seq)\n                for chain in current_chains:\n                    sequences[chain] = seq\n            header = line[1:]\n            chain_match = re.search(r'Chains?\\s+([^\\|]+)', header)\n            current_chains = []\n            if chain_match:\n                for std_chain, auth_chain in re.findall(\n                    r'([A-Za-z0-9]+)(?:\\[auth\\s+([A-Za-z0-9]+)\\])?', chain_match.group(1)\n                ):\n                    if std_chain:\n                        current_chains.append(std_chain)\n                        auth_to_std[auth_chain if auth_chain else std_chain] = std_chain\n            current_seq = []\n        else:\n            current_seq.append(line)\n    if current_chains and current_seq:\n        seq = ''.join(current_seq)\n        for chain in current_chains:\n            sequences[chain] = seq\n    return sequences, auth_to_std\n\n\ndef parse_stoichiometry(stoich_str):\n    if pd.isna(stoich_str) or str(stoich_str).strip() == \"\":\n        return []\n    result = []\n    for part in str(stoich_str).split(';'):\n        part = part.strip()\n        m = re.match(r'([^:]+):(\\d+)', part)\n        if m:\n            result.append((m.group(1).strip(), int(m.group(2))))\n    return result\n\n\ndef get_chain_sequences(stoich_str, all_seq_str, full_sequence):\n    stoichiometry = parse_stoichiometry(stoich_str)\n    sequences, auth_to_std = parse_fasta_with_auth(all_seq_str)\n    chain_data, idx = [], 0\n    for auth_id, count in stoichiometry:\n        std_id = auth_to_std.get(auth_id, auth_id)\n        if std_id not in sequences:\n            continue\n        cseq = sequences[std_id]\n        clen = len(cseq)\n        for copy in range(count):\n            chain_data.append({\n                'chain_id': auth_id,\n                'copy_num': copy + 1,\n                'start_idx': idx,\n                'end_idx': idx + clen,\n                'sequence': full_sequence[idx:idx + clen],\n            })\n            idx += clen\n    return chain_data\n\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:\n            continue\n        if line.startswith(\">\"):\n            if cur is not None:\n                out[cur] = \"\".join(parts)\n            cur = line[1:].split()[0]\n            parts = []\n        else:\n            parts.append(line.replace(\" \", \"\"))\n    if cur is not None:\n        out[cur] = \"\".join(parts)\n    return out\n\n\ndef get_chain_segments(row) -> list:\n    seq = row[\"sequence\"]\n    stoich, all_sq = row.get(\"stoichiometry\", \"\"), row.get(\"all_sequences\", \"\")\n    if pd.isna(stoich) or pd.isna(all_sq) or not str(stoich).strip() or not str(all_sq).strip():\n        return [(0, len(seq))]\n    try:\n        cd = get_chain_sequences(stoich, all_sq, seq)\n        if not cd:\n            return [(0, len(seq))]\n        segs = [(c['start_idx'], c['end_idx']) for c in cd]\n        return segs if segs[-1][1] == len(seq) else [(0, len(seq))]\n    except Exception:\n        return [(0, len(seq))]\n\n\ndef get_query_chains(row) -> list:\n    seq = row[\"sequence\"]\n    stoich_raw = row.get(\"stoichiometry\", \"\")\n    all_seq_raw = row.get(\"all_sequences\", \"\")\n    if pd.isna(stoich_raw) or pd.isna(all_seq_raw) or not str(stoich_raw).strip() or not str(all_seq_raw).strip():\n        return [{\"chain_id\": \"A\", \"sequence\": seq, \"start\": 0, \"end\": len(seq), \"copy\": 1}]\n    try:\n        cd = get_chain_sequences(stoich_raw, all_seq_raw, seq)\n        if not cd:\n            return [{\"chain_id\": \"A\", \"sequence\": seq, \"start\": 0, \"end\": len(seq), \"copy\": 1}]\n        result = [{\n            \"chain_id\": c[\"chain_id\"],\n            \"sequence\": c[\"sequence\"],\n            \"start\": c[\"start_idx\"],\n            \"end\": c[\"end_idx\"],\n            \"copy\": c[\"copy_num\"]\n        } for c in cd]\n        if result[-1][\"end\"] != len(seq):\n            return [{\"chain_id\": \"A\", \"sequence\": seq, \"start\": 0, \"end\": len(seq), \"copy\": 1}]\n        return result\n    except Exception:\n        return [{\"chain_id\": \"A\", \"sequence\": seq, \"start\": 0, \"end\": len(seq), \"copy\": 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        raw_s = r.get(\"stoichiometry\", \"\")\n        stoich_map[tid] = \"\" if pd.isna(raw_s) else str(raw_s)\n    return seg_map, stoich_map\n\n\ndef build_training_database(train_seqs_df, train_labels_df):\n    \"\"\"Build full_coords dict and chain_index dict from training data.\"\"\"\n    print(\"  Building training database...\")\n    t0 = time.time()\n\n    full_coords = {}\n    train_labels_df[\"target_id\"] = train_labels_df[\"ID\"].str.rsplit(\"_\", n=1).str[0]\n    for tid, grp in train_labels_df.groupby(\"target_id\", sort=False):\n        g = grp.sort_values(\"resid\")\n        arr = np.column_stack([\n            pd.to_numeric(g[\"x_1\"], errors=\"coerce\").values,\n            pd.to_numeric(g[\"y_1\"], errors=\"coerce\").values,\n            pd.to_numeric(g[\"z_1\"], errors=\"coerce\").values,\n        ])\n        full_coords[tid] = np.nan_to_num(arr, nan=0.0)\n\n    chain_index = {}\n    n_with_chains = 0\n    for _, row in train_seqs_df.iterrows():\n        tid = row[\"target_id\"]\n        if tid not in full_coords:\n            continue\n        full_seq = row[\"sequence\"]\n        stoich_raw = row.get(\"stoichiometry\", \"\")\n        all_seq_raw = row.get(\"all_sequences\", \"\")\n\n        chains = []\n        if pd.notna(stoich_raw) and pd.notna(all_seq_raw) and str(stoich_raw).strip():\n            cd = get_chain_sequences(stoich_raw, all_seq_raw, full_seq)\n            if cd and cd[-1][\"end_idx\"] == len(full_seq):\n                rna_chars = set(\"ACGU\")\n                for c in cd:\n                    non_rna = sum(1 for ch in c[\"sequence\"] if ch not in rna_chars)\n                    chains.append({\n                        \"chain_id\": c[\"chain_id\"],\n                        \"sequence\": c[\"sequence\"],\n                        \"start\": c[\"start_idx\"],\n                        \"end\": c[\"end_idx\"],\n                        \"is_rna\": len(c[\"sequence\"]) > 0 and non_rna / len(c[\"sequence\"]) <= 0.2,\n                    })\n                n_with_chains += 1\n\n        if not chains:\n            chains = [{\n                \"chain_id\": \"A\",\n                \"sequence\": full_seq,\n                \"start\": 0,\n                \"end\": len(full_seq),\n                \"is_rna\": True\n            }]\n\n        chain_index[tid] = {\"full_seq\": full_seq, \"chains\": chains}\n\n    print(f\"  DB built: {len(full_coords)} entries, \"\n          f\"{n_with_chains} with chains ({time.time()-t0:.1f}s)\")\n    return full_coords, chain_index\n\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","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ═══════════════════════════════════════════════════════════════════════════════\n# §8  Template Search & Adaptation  (FULL REPLACEMENT)\n# ═══════════════════════════════════════════════════════════════════════════════\n\ndef _align_pair(seq1, seq2):\n    \"\"\"Align two sequences, return (normalized_score, percent_identity).\"\"\"\n    aln = next(iter(_aligner.align(seq1, seq2)))\n    norm_s = aln.score / (2 * min(len(seq1), len(seq2)))\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 seq1[qp] == seq2[tp]\n    )\n    return norm_s, 100 * identical / len(seq1)\n\n\ndef _fill_gaps(coords):\n    \"\"\"Fill NaN gaps by interpolation/extrapolation.\"\"\"\n    for i in range(len(coords)):\n        if np.isnan(coords[i, 0]):\n            pv = next((j for j in range(i - 1, -1, -1) if not np.isnan(coords[j, 0])), -1)\n            nv = next((j for j in range(i + 1, len(coords)) if not np.isnan(coords[j, 0])), -1)\n            if pv >= 0 and nv >= 0:\n                w = (i - pv) / (nv - pv)\n                coords[i] = (1 - w) * coords[pv] + w * coords[nv]\n            elif pv >= 0:\n                coords[i] = coords[pv] + [3, 0, 0]\n            elif nv >= 0:\n                coords[i] = coords[nv] + [3, 0, 0]\n            else:\n                coords[i] = [i * 3, 0, 0]\n    return coords\n\n\ndef _align_and_map(query_seq, template_seq, template_coords):\n    \"\"\"Align two sequences and map coords. Returns array with NaN for gaps.\"\"\"\n    aln = next(iter(_aligner.align(query_seq, template_seq)))\n    new_coords = np.full((len(query_seq), 3), np.nan)\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    return new_coords\n\n\ndef adapt_template_full(query_seq, template_seq, template_coords):\n    \"\"\"Full-sequence adaptation.\"\"\"\n    coords = _align_and_map(query_seq, template_seq, template_coords)\n    return np.nan_to_num(_fill_gaps(coords))\n\n\ndef adapt_template_chain_aware(query_chains, chain_matches, template_full_coords):\n    \"\"\"\n    Chain-aware adaptation: for each query chain instance, align against ONLY\n    the matched training chain's region.\n    \"\"\"\n    total_len = query_chains[-1][\"end\"]\n    new_coords = np.full((total_len, 3), np.nan)\n\n    for qc in query_chains:\n        qkey = query_chain_key(qc)\n        qseq = qc[\"sequence\"]\n        q_start, q_end = qc[\"start\"], qc[\"end\"]\n\n        tc = chain_matches.get(qkey)\n        if tc is None:\n            continue\n\n        tc_seq = tc[\"sequence\"]\n        tc_coords = template_full_coords[tc[\"start\"]:tc[\"end\"]]\n\n        if len(tc_coords) != len(tc_seq):\n            continue\n\n        mapped = _align_and_map(qseq, tc_seq, tc_coords)\n        new_coords[q_start:q_end] = mapped\n\n    return np.nan_to_num(_fill_gaps(new_coords))\n\n\n# ── Candidate Scoring ─────────────────────────────────────────────────────────\n\ndef score_entry_full(train_entry, full_query_seq):\n    \"\"\"Returns a FULL candidate dict or None.\"\"\"\n    train_full_seq = train_entry[\"full_seq\"]\n    lr = abs(len(train_full_seq) - len(full_query_seq)) / max(len(train_full_seq), len(full_query_seq))\n    if lr > FULL_LENGTH_RATIO:\n        return None\n    sim, pct = _align_pair(full_query_seq, train_full_seq)\n    return {\n        \"norm_score\": sim,\n        \"avg_pct\": pct,\n        \"min_pct\": pct,\n        \"strategy\": \"full\",\n        \"chain_matches\": {},\n    }\n\n\ndef score_entry_chain(query_chains, train_entry):\n    \"\"\"Returns a CHAIN candidate dict or None.\"\"\"\n    train_chains = train_entry[\"chains\"]\n    chain_scores, chain_pcts, chain_matches = [], [], {}\n\n    for qc in query_chains:\n        qkey = query_chain_key(qc)\n        qseq = qc[\"sequence\"]\n\n        best_sim, best_pct, best_tc = -1.0, 0.0, None\n        for tc in train_chains:\n            if not tc[\"is_rna\"] or len(tc[\"sequence\"]) < 5:\n                continue\n            if abs(len(tc[\"sequence\"]) - len(qseq)) / max(len(tc[\"sequence\"]), len(qseq)) > CHAIN_LENGTH_RATIO:\n                continue\n            sim, pct = _align_pair(qseq, tc[\"sequence\"])\n            if sim > best_sim:\n                best_sim, best_pct, best_tc = sim, pct, tc\n\n        if best_sim >= 0:\n            chain_scores.append(best_sim)\n            chain_pcts.append(best_pct)\n            chain_matches[qkey] = best_tc\n\n    if not chain_scores or len(chain_scores) != len(query_chains):\n        return None\n\n    chain_avg_sim = sum(chain_scores) / len(chain_scores)\n    chain_avg_pct = sum(chain_pcts) / len(chain_pcts)\n    chain_min_pct = min(chain_pcts)\n\n    if chain_min_pct < CHAIN_TRUST_MIN_PCT:\n        return None\n\n    return {\n        \"norm_score\": chain_avg_sim,\n        \"avg_pct\": chain_avg_pct,\n        \"min_pct\": chain_min_pct,\n        \"strategy\": \"chain\",\n        \"chain_matches\": chain_matches,\n    }\n\n\ndef generate_entry_candidates(query_chains, full_query_seq, tid, entry, full_coords):\n    \"\"\"\n    Each training entry may contribute one FULL and one CHAIN candidate.\n    Both are emitted if they pass thresholds.\n    \"\"\"\n    candidates = []\n\n    full_cand = score_entry_full(entry, full_query_seq)\n    if full_cand is not None:\n        if full_cand[\"norm_score\"] >= MIN_SIMILARITY and full_cand[\"avg_pct\"] >= MIN_PERCENT_IDENTITY:\n            candidates.append({\n                \"pdb_id\": tid,\n                \"norm_score\": full_cand[\"norm_score\"],\n                \"avg_pct\": full_cand[\"avg_pct\"],\n                \"min_pct\": full_cand[\"min_pct\"],\n                \"strategy\": full_cand[\"strategy\"],\n                \"chain_matches\": full_cand[\"chain_matches\"],\n                \"full_seq\": entry[\"full_seq\"],\n                \"full_coords\": full_coords[tid],\n            })\n\n    chain_cand = score_entry_chain(query_chains, entry)\n    if chain_cand is not None:\n        if chain_cand[\"norm_score\"] >= MIN_SIMILARITY and chain_cand[\"avg_pct\"] >= MIN_PERCENT_IDENTITY:\n            candidates.append({\n                \"pdb_id\": tid,\n                \"norm_score\": chain_cand[\"norm_score\"],\n                \"avg_pct\": chain_cand[\"avg_pct\"],\n                \"min_pct\": chain_cand[\"min_pct\"],\n                \"strategy\": chain_cand[\"strategy\"],\n                \"chain_matches\": chain_cand[\"chain_matches\"],\n                \"full_seq\": entry[\"full_seq\"],\n                \"full_coords\": full_coords[tid],\n            })\n\n    return candidates\n\n\ndef find_best_entries(query_chains, full_query_seq, full_coords, chain_index, top_n=30):\n    \"\"\"Pool FULL and CHAIN candidates from all training entries, rank together.\"\"\"\n    results = []\n    for tid, entry in chain_index.items():\n        results.extend(\n            generate_entry_candidates(query_chains, full_query_seq, tid, entry, full_coords)\n        )\n    results.sort(key=lambda x: x[\"norm_score\"], reverse=True)\n    return results[:top_n]","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ═══════════════════════════════════════════════════════════════════════════════\n# §9  Structure Refinement & Diversity Transforms\n# ═══════════════════════════════════════════════════════════════════════════════\n\ndef adaptive_rna_constraints(coords, target_id, segments_map, confidence=1.0, passes=2) -> np.ndarray:\n    X = coords.copy()\n    segments = segments_map.get(target_id, [(0, len(X))])\n    strength = max(0.75 * (1.0 - min(confidence, 0.97)), 0.02)\n    for _ in range(passes):\n        for s, e in segments:\n            C = X[s:e]\n            L = e - s\n            if L < 3:\n                continue\n            # bond i–i+1 ~5.95 Å\n            d = C[1:] - C[:-1]\n            dist = np.linalg.norm(d, axis=1) + 1e-6\n            adj = d * ((5.95 - dist) / dist)[:, None] * (0.22 * strength)\n            C[:-1] -= adj\n            C[1:] += adj\n            # soft i–i+2 ~10.2 Å\n            d2 = C[2:] - C[:-2]\n            d2n = np.linalg.norm(d2, axis=1) + 1e-6\n            adj2 = d2 * ((10.2 - d2n) / d2n)[:, None] * (0.10 * strength)\n            C[:-2] -= adj2\n            C[2:] += adj2\n            # Laplacian smoothing\n            C[1:-1] += (0.06 * strength) * (0.5 * (C[:-2] + C[2:]) - C[1:-1])\n            # Self-avoidance\n            if L >= 25:\n                idx = (np.linspace(0, L - 1, min(L, 160)).astype(int)\n                       if L > 220 else np.arange(L))\n                P = C[idx]\n                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    return X\n\n\ndef _rotmat(axis, ang):\n    a = np.asarray(axis, float)\n    a /= np.linalg.norm(a) + 1e-12\n    x, y, z = a\n    c, s = np.cos(ang), np.sin(ang)\n    CC = 1 - c\n    return np.array([\n        [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    ])\n\n\ndef apply_hinge(coords, seg, rng, deg=22):\n    s, e = seg\n    L = e - s\n    if L < 30:\n        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()\n    p0 = X[pivot].copy()\n    X[pivot + 1:e] = (X[pivot + 1:e] - p0) @ R.T + p0\n    return X\n\n\ndef jitter_chains(coords, segs, rng, deg=12, trans=1.5):\n    X = coords.copy()\n    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)\n        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\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:\n            continue\n        ctrl = np.linspace(0, L - 1, 6)\n        disp = rng.normal(0, amp, (6, 3))\n        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\ndef generate_rna_structure(sequence: str, seed=None) -> np.ndarray:\n    \"\"\"Idealized A-form RNA helix — last-resort de-novo fallback.\"\"\"\n    if seed is not None:\n        np.random.seed(seed)\n    n = len(sequence)\n    coords = np.zeros((n, 3))\n    for i in range(n):\n        ang = i * 0.6\n        coords[i] = [10.0 * np.cos(ang), 10.0 * np.sin(ang), i * 2.5]\n    return coords","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ═══════════════════════════════════════════════════════════════════════════════\n# §9.5  TM-Score & Diversity Selection\n# ═══════════════════════════════════════════════════════════════════════════════\n\ndef _kabsch_superpose(mobile: np.ndarray, target: np.ndarray) -> np.ndarray:\n    \"\"\"Kabsch-align mobile onto target, return transformed mobile.\"\"\"\n    # Filter out zero/nan rows for fitting\n    valid = (\n        np.isfinite(mobile).all(axis=1)\n        & np.isfinite(target).all(axis=1)\n        & (np.abs(mobile).sum(axis=1) > 1e-6)\n        & (np.abs(target).sum(axis=1) > 1e-6)\n    )\n    if valid.sum() < 3:\n        return mobile.copy()\n\n    cm = mobile[valid].mean(axis=0)\n    ct = target[valid].mean(axis=0)\n    m_c = mobile[valid] - cm\n    t_c = target[valid] - ct\n\n    H = m_c.T @ t_c\n    U, S, Vt = np.linalg.svd(H)\n    d = np.linalg.det(Vt.T @ U.T)\n    sign = np.diag([1.0, 1.0, np.sign(d)])\n    R = Vt.T @ sign @ U.T\n\n    return (mobile - cm) @ R.T + ct\n\n\ndef tm_score_rna(coords1: np.ndarray, coords2: np.ndarray,\n                 superpose: bool = True) -> float:\n    \"\"\"\n    Compute TM-score between two RNA C1' coordinate arrays.\n\n    Uses the RNA-specific d0: d0 = 0.6 * sqrt(L - 0.5) - 2.5\n    Normalized by the length of coords2 (the reference).\n\n    Parameters\n    ----------\n    coords1, coords2 : (L, 3) arrays\n    superpose : if True, Kabsch-align coords1 onto coords2 first\n\n    Returns\n    -------\n    tm : float in [0, 1], higher = more similar\n    \"\"\"\n    L = len(coords2)\n    if len(coords1) != L:\n        minL = min(len(coords1), L)\n        coords1 = coords1[:minL]\n        coords2 = coords2[:minL]\n        L = minL\n    if L < 3:\n        return 1.0\n\n    d0 = max(0.6 * np.sqrt(L - 0.5) - 2.5, 0.5)\n\n    if superpose:\n        c1 = _kabsch_superpose(coords1.copy(), coords2.copy())\n    else:\n        c1 = coords1\n\n    dist = np.sqrt(np.sum((c1 - coords2) ** 2, axis=1))\n    tm = np.sum(1.0 / (1.0 + (dist / d0) ** 2)) / L\n    return float(tm)\n\n\ndef pairwise_tm_matrix(coords_list: list) -> np.ndarray:\n    \"\"\"\n    Pairwise TM-score matrix.  Uses max(TM(a→b), TM(b→a)) for symmetry.\n    \"\"\"\n    N = len(coords_list)\n    tm_mat = np.eye(N, dtype=np.float64)\n    for i in range(N):\n        for j in range(i + 1, N):\n            tm_ij = tm_score_rna(coords_list[i], coords_list[j])\n            tm_ji = tm_score_rna(coords_list[j], coords_list[i])\n            tm_sym = max(tm_ij, tm_ji)\n            tm_mat[i, j] = tm_mat[j, i] = tm_sym\n    return tm_mat\n\n\ndef pool_diversity_stats(coords_list: list) -> tuple:\n    \"\"\"\n    Returns (mean_pairwise_TM, max_pairwise_TM) for a pool of structures.\n    Lower values = more diverse.\n    \"\"\"\n    if len(coords_list) < 2:\n        return 1.0, 1.0\n    tm_mat = pairwise_tm_matrix(coords_list)\n    N = len(coords_list)\n    off_diag = tm_mat[np.triu_indices(N, k=1)]\n    return float(np.mean(off_diag)), float(np.max(off_diag))\n\n\ndef select_diverse_from_pool(coords_pool: list, scores: list,\n                             n_select: int = 5,\n                             quality_weight: float = 0.4) -> list:\n    \"\"\"\n    Greedy max-diversity selection with quality trade-off.\n\n    1. Start with the highest-scoring sample.\n    2. Iteratively add the sample that maximizes:\n         quality_weight * normalized_score + (1 - quality_weight) * (1 - max_TM_to_selected)\n\n    Parameters\n    ----------\n    coords_pool  : list of (L, 3) arrays\n    scores       : list/array of confidence scores (higher = better)\n    n_select     : how many to pick\n    quality_weight : 0 = pure diversity, 1 = pure quality\n\n    Returns\n    -------\n    list of selected indices into coords_pool\n    \"\"\"\n    M = len(coords_pool)\n    if M <= n_select:\n        return list(range(M))\n\n    scores = np.array(scores, dtype=np.float64)\n    s_range = scores.max() - scores.min()\n    if s_range > 1e-8:\n        s_norm = (scores - scores.min()) / s_range\n    else:\n        s_norm = np.ones(M) * 0.5\n\n    # Precompute pairwise TM\n    tm_mat = pairwise_tm_matrix(coords_pool)\n\n    selected = [int(np.argmax(scores))]\n\n    for _ in range(n_select - 1):\n        best_idx, best_val = -1, -np.inf\n        for c in range(M):\n            if c in selected:\n                continue\n            max_tm = max(tm_mat[c, s] for s in selected)\n            diversity = 1.0 - max_tm\n            val = quality_weight * s_norm[c] + (1.0 - quality_weight) * diversity\n            if val > best_val:\n                best_val = val\n                best_idx = c\n        if best_idx < 0:\n            break\n        selected.append(best_idx)\n\n    return selected","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ═══════════════════════════════════════════════════════════════════════════════\n# §10  Phase 1 — Template-Based Modeling  (FULL REPLACEMENT)\n#      Now uses pooled FULL + CHAIN candidate routing from Script 2\n# ═══════════════════════════════════════════════════════════════════════════════\n\ndef tbm_phase(test_df, full_coords, chain_index, segments_map):\n    \"\"\"\n    Returns\n    -------\n    template_predictions : {target_id: [np.ndarray(seq_len, 3), ...]}\n    protenix_queue       : {target_id: (n_model_slots, full_sequence)}\n    boltz_queue          : {target_id: (1, full_sequence)}\n    \"\"\"\n    print(f\"\\n{'=' * 60}\")\n    print(f\"PHASE 1: TBM v5 — Pooled FULL + CHAIN Candidate Routing\")\n    print(f\"  MIN_SIMILARITY       = {MIN_SIMILARITY}\")\n    print(f\"  MIN_PCT_IDENTITY     = {MIN_PERCENT_IDENTITY}\")\n    print(f\"  CHAIN_TRUST_MIN_PCT  = {CHAIN_TRUST_MIN_PCT}\")\n    print(f\"  FULL_LENGTH_RATIO    = {FULL_LENGTH_RATIO}\")\n    print(f\"  CHAIN_LENGTH_RATIO   = {CHAIN_LENGTH_RATIO}\")\n    print(f\"  Short (≤{MAX_SEQ_LEN}nt): max_tbm=2, min_ptx=3\")\n    print(f\"  Long  (>{MAX_SEQ_LEN}nt): max_tbm={N_SAMPLE}, min_ptx=0  (TBM-first)\")\n    print(f\"  Boltz2 ensembling for targets < {BOLTZ_MAX_TOKENS} nt\")\n    print(f\"  Training entries     = {len(chain_index)}\")\n    print(f\"{'=' * 60}\")\n    t0 = time.time()\n\n    template_predictions: dict = {}\n    protenix_queue: dict = {}\n    boltz_queue: dict = {}\n    stats = {\"chain\": 0, \"full\": 0, \"none\": 0}\n\n    for _, row in tqdm(test_df.iterrows(), total=len(test_df), desc=\"TBM\"):\n        tid = row[\"target_id\"]\n        full_seq = row[\"sequence\"]\n        full_len = len(full_seq)\n        segs = segments_map.get(tid, [(0, full_len)])\n        query_chains = get_query_chains(row)\n        n_chains = len(query_chains)\n        unique_chain_seqs = len(set(qc[\"sequence\"] for qc in query_chains))\n\n        # ── Length-adaptive slot allocation ────────────────────────────────\n        max_tbm_slots, min_ptx_slots = get_slot_allocation(full_len)\n        is_long = full_len > MAX_SEQ_LEN\n\n        # ── Find pooled FULL + CHAIN candidates ──────────────────────────\n        best_entries = find_best_entries(\n            query_chains, full_seq, full_coords, chain_index, top_n=30\n        )\n\n        preds, used = [], set()\n        candidates_with_scores = []\n\n        for i, entry in enumerate(best_entries):\n            if len(preds) >= N_SAMPLE:\n                break\n\n            # Deduplicate by (template, strategy)\n            use_key = (entry[\"pdb_id\"], entry[\"strategy\"])\n            if use_key in used:\n                continue\n\n            if entry[\"strategy\"] == \"chain\" and entry[\"chain_matches\"]:\n                adapted = adapt_template_chain_aware(\n                    query_chains, entry[\"chain_matches\"], entry[\"full_coords\"]\n                )\n            else:\n                adapted = adapt_template_full(\n                    full_seq, entry[\"full_seq\"], entry[\"full_coords\"]\n                )\n\n            rng = np.random.default_rng((stable_hash(tid) + i * 10007) % (2**32))\n            slot = len(preds)\n            if slot == 0:\n                X = adapted\n            elif slot == 1:\n                X = adapted + rng.normal(0, 0.01, adapted.shape)\n            elif slot == 2:\n                longest = max(segs, key=lambda se: se[1] - se[0])\n                X = apply_hinge(adapted, longest, rng)\n            elif slot == 3:\n                X = jitter_chains(adapted, segs, rng)\n            else:\n                X = smooth_wiggle(adapted, segs, rng)\n\n            refined = adaptive_rna_constraints(\n                X, tid, segments_map, confidence=entry[\"norm_score\"]\n            )\n            candidates_with_scores.append((refined, entry[\"norm_score\"]))\n            preds.append(refined)\n            used.add(use_key)\n\n        # ── Diversity-select TBM ──────────────────────────────────────────\n        if len(candidates_with_scores) > max_tbm_slots:\n            coords_list = [c for c, _ in candidates_with_scores]\n            scores_list = [s for _, s in candidates_with_scores]\n            indices = select_diverse_from_pool(\n                coords_list, scores_list, max_tbm_slots, QUALITY_WEIGHT\n            )\n            preds = [candidates_with_scores[i][0] for i in indices]\n            print(f\"  {tid} ({full_len} nt{'*' if is_long else ''}, \"\n                  f\"{n_chains}ch/{unique_chain_seqs}uniq): \"\n                  f\"{len(candidates_with_scores)} TBM candidates \"\n                  f\"→ selected best {len(preds)} by diversity\")\n        else:\n            preds = preds[:max_tbm_slots]\n\n        template_predictions[tid] = preds\n\n        # ── Log top entry info ────────────────────────────────────────────\n        if best_entries:\n            top = best_entries[0]\n            stats[top[\"strategy\"]] += 1\n        else:\n            stats[\"none\"] += 1\n\n        # ── Remaining model slots ─────────────────────────────────────────\n        n_model_slots = N_SAMPLE - len(preds)\n        tag = \"LONG/TBM-first\" if is_long else \"short\"\n\n        if len(full_seq) < BOLTZ_MAX_TOKENS and n_model_slots > 0:\n            boltz_queue[tid] = (1, full_seq)\n            n_protenix = n_model_slots - 1\n            if n_protenix > 0:\n                protenix_queue[tid] = (n_protenix, full_seq)\n            if best_entries:\n                top = best_entries[0]\n                print(f\"  {tid} ({full_len} nt, {tag}, \"\n                      f\"{n_chains}ch/{unique_chain_seqs}uniq): \"\n                      f\"top={top['pdb_id']} sim={top['norm_score']:.3f} \"\n                      f\"pct={top['avg_pct']:.1f}% min_pct={top['min_pct']:.1f}% \"\n                      f\"via {top['strategy']} | \"\n                      f\"{len(preds)} TBM → {n_protenix} Protenix + 1 Boltz2\")\n            else:\n                print(f\"  {tid} ({full_len} nt, {tag}): no templates \"\n                      f\"→ {n_protenix} Protenix + 1 Boltz2\")\n        elif n_model_slots > 0:\n            protenix_queue[tid] = (n_model_slots, full_seq)\n            if best_entries:\n                top = best_entries[0]\n                print(f\"  {tid} ({full_len} nt, {tag}, \"\n                      f\"{n_chains}ch/{unique_chain_seqs}uniq): \"\n                      f\"top={top['pdb_id']} sim={top['norm_score']:.3f} \"\n                      f\"pct={top['avg_pct']:.1f}% min_pct={top['min_pct']:.1f}% \"\n                      f\"via {top['strategy']} | \"\n                      f\"{len(preds)} TBM → {n_model_slots} Protenix \"\n                      f\"(too long for Boltz2)\")\n            else:\n                print(f\"  {tid} ({full_len} nt, {tag}): no templates \"\n                      f\"→ {n_model_slots} Protenix\")\n        else:\n            if best_entries:\n                top = best_entries[0]\n                print(f\"  {tid} ({full_len} nt, {tag}, \"\n                      f\"{n_chains}ch/{unique_chain_seqs}uniq): \"\n                      f\"top={top['pdb_id']} sim={top['norm_score']:.3f} \"\n                      f\"pct={top['avg_pct']:.1f}% min_pct={top['min_pct']:.1f}% \"\n                      f\"via {top['strategy']} | \"\n                      f\"{len(preds)} TBM → fully covered\")\n            else:\n                print(f\"  {tid} ({full_len} nt, {tag}): fully covered (de-novo)\")\n\n    elapsed = time.time() - t0\n    print(f\"\\nPhase 1 done in {elapsed:.1f}s\")\n    print(f\"  Strategy stats: chain={stats['chain']}, \"\n          f\"full={stats['full']}, none={stats['none']}\")\n    return template_predictions, protenix_queue, boltz_queue","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ═══════════════════════════════════════════════════════════════════════════════\n# §11  Phase 1.5 — RNAPro Refinement   (COMPLETE REPLACEMENT)\n# ═══════════════════════════════════════════════════════════════════════════════\n\nimport subprocess as _subprocess\n\ndef _write_rnapro_runner_code():\n    \"\"\"Write the RNAPro inference runner Python script (multi-GPU aware).\"\"\"\n\n    rnapro_runner_code = r\"\"\"\nimport os, shutil, logging, traceback, warnings, argparse, json\nfrom contextlib import nullcontext\nfrom os.path import join as opjoin\nimport torch\nimport pandas as pd\nimport numpy as np\nfrom biotite.structure.io import pdbx\n\nfrom configs.configs_base import configs as configs_base\nfrom configs.configs_data import data_configs\nfrom configs.configs_inference import inference_configs\nfrom runner.dumper import DataDumper\nfrom rnapro.config import parse_sys_args\nfrom rnapro.config.config import ConfigManager, ArgumentNotSet\nfrom rnapro.data.infer_data_pipeline import get_inference_dataloader\nfrom rnapro.model.RNAPro import RNAPro\nfrom rnapro.utils.distributed import DIST_WRAPPER\nfrom rnapro.utils.seed import seed_everything\nfrom rnapro.utils.torch_utils import to_device\n\nwarnings.filterwarnings(\"ignore\", category=UserWarning)\nwarnings.filterwarnings(\"ignore\", category=FutureWarning)\nwarnings.filterwarnings(\"ignore\", category=DeprecationWarning)\nlogging.basicConfig(level=logging.WARNING)\nlogging.getLogger(\"rnapro\").setLevel(logging.WARNING)\n\ndef parse_configs(configs, arg_str=None, fill_required_with_null=False):\n    manager = ConfigManager(configs, fill_required_with_null=fill_required_with_null)\n    parser = argparse.ArgumentParser()\n    parser.add_argument(\"--max_len\", type=int, default=10000, required=False)\n    for key, (dtype, default_value, allow_none, required) in manager.config_infos.items():\n        parser.add_argument(\"--\" + key, type=str, default=ArgumentNotSet(), required=required)\n    merged_configs = manager.merge_configs(\n        vars(parser.parse_args(arg_str.split())) if arg_str else {})\n    merged_configs.max_len = parser.parse_args(arg_str.split()).max_len\n    return merged_configs\n\nclass dotdict(dict):\n    __setattr__ = dict.__setitem__; __delattr__ = dict.__delitem__\n    def __getattr__(self, name):\n        try: return self[name]\n        except KeyError: raise AttributeError(name)\n\nclass InferenceRunner:\n    def __init__(self, configs):\n        self.configs = configs\n        self.init_env(); self.init_basics(); self.init_model()\n        self.load_checkpoint()\n        self.init_dumper(need_atom_confidence=configs.need_atom_confidence,\n                         sorted_by_ranking_score=configs.sorted_by_ranking_score)\n    def init_env(self):\n        self.use_cuda = torch.cuda.device_count() > 0\n        self.device = torch.device(\"cuda:{}\".format(DIST_WRAPPER.local_rank) if self.use_cuda else \"cpu\")\n        if self.use_cuda: torch.cuda.set_device(self.device)\n    def init_basics(self):\n        self.dump_dir = self.configs.dump_dir; self.error_dir = opjoin(self.dump_dir, \"ERR\")\n        os.makedirs(self.dump_dir, exist_ok=True); os.makedirs(self.error_dir, exist_ok=True)\n    def init_model(self):\n        self.model = RNAPro(self.configs).to(self.device)\n    def load_checkpoint(self):\n        ckpt_path = self.configs.load_checkpoint_path\n        if not os.path.exists(ckpt_path): raise Exception(f\"Checkpoint missing: {ckpt_path}\")\n        ckpt = torch.load(ckpt_path, self.device)\n        if list(ckpt[\"model\"].keys())[0].startswith(\"module.\"):\n            ckpt[\"model\"] = {k[7:]: v for k,v in ckpt[\"model\"].items()}\n        self.model.load_state_dict(state_dict=ckpt[\"model\"], strict=True)\n        self.model.eval()\n    def init_dumper(self, need_atom_confidence=False, sorted_by_ranking_score=True):\n        self.dumper = DataDumper(base_dir=self.dump_dir,\n                                 need_atom_confidence=need_atom_confidence,\n                                 sorted_by_ranking_score=sorted_by_ranking_score)\n    @torch.no_grad()\n    def predict(self, data):\n        prec = {\"fp32\": torch.float32, \"bf16\": torch.bfloat16, \"fp16\": torch.float16}[self.configs.dtype]\n        ctx = torch.autocast(device_type=\"cuda\", dtype=prec) if torch.cuda.is_available() else nullcontext()\n        data = to_device(data, self.device)\n        with ctx:\n            prediction, _, _ = self.model(input_feature_dict=data[\"input_feature_dict\"],\n                                          label_full_dict=None, label_dict=None, mode=\"inference\")\n        return prediction\n    def update_model_configs(self, new_configs):\n        self.model.configs = new_configs\n\ndef update_inference_configs(configs, N_token):\n    if N_token > 3840:   configs.skip_amp.confidence_head = False; configs.skip_amp.sample_diffusion = False\n    elif N_token > 2560: configs.skip_amp.confidence_head = False; configs.skip_amp.sample_diffusion = True\n    else:                configs.skip_amp.confidence_head = True;  configs.skip_amp.sample_diffusion = True\n    return configs\n\ndef infer_predict(runner, configs):\n    try: dataloader = get_inference_dataloader(configs=configs)\n    except Exception as e:\n        with open(opjoin(runner.error_dir, \"error.txt\"), \"a\") as f: f.write(f\"{e}\\n{traceback.format_exc()}\")\n        return\n    for seed in configs.seeds:\n        seed_everything(seed=seed, deterministic=configs.deterministic)\n        for batch in dataloader:\n            try:\n                data, atom_array, data_error_message = batch[0]\n                sample_name = data[\"sample_name\"]\n                if len(data_error_message) > 0:\n                    with open(opjoin(runner.error_dir, f\"{sample_name}.txt\"), \"a\") as f: f.write(data_error_message)\n                    continue\n                new_configs = update_inference_configs(configs, data[\"N_token\"].item())\n                runner.update_model_configs(new_configs)\n                prediction = runner.predict(data)\n                runner.dumper.dump(dataset_name=\"\", pdb_id=sample_name, seed=seed,\n                                   pred_dict=prediction, atom_array=atom_array,\n                                   entity_poly_type=data[\"entity_poly_type\"])\n                torch.cuda.empty_cache()\n            except Exception as e:\n                with open(opjoin(runner.error_dir, f\"{sample_name}.txt\"), \"a\") as f:\n                    f.write(f\"{e}\\n{traceback.format_exc()}\")\n                if hasattr(torch.cuda, \"empty_cache\"): torch.cuda.empty_cache()\n\ndef make_dummy_solution(df):\n    sol = dotdict()\n    for _, r in df.iterrows():\n        sol[r.target_id] = dotdict(target_id=r.target_id, sequence=r.sequence, coord=[])\n    return sol\n\ndef solution_to_submit_df(solution):\n    frames = []\n    for k, s in solution.items():\n        L = len(s.sequence)\n        df = pd.DataFrame()\n        df[\"ID\"] = [f\"{s.target_id}_{i+1}\" for i in range(L)]\n        df[\"resname\"] = list(s.sequence); df[\"resid\"] = list(range(1, L+1))\n        for j, c in enumerate(s.coord):\n            df[f\"x_{j+1}\"] = c[:,0]; df[f\"y_{j+1}\"] = c[:,1]; df[f\"z_{j+1}\"] = c[:,2]\n        frames.append(df)\n    return pd.concat(frames)\n\ndef extract_c1_coordinates(cif_file_path):\n    from biotite.structure.io import pdbx\n    try:\n        with open(cif_file_path, \"r\") as f: cif_data = pdbx.CIFFile.read(f)\n        atom_array = pdbx.get_structure(cif_data, model=1)\n        mask  = np.char.strip(atom_array.atom_name.astype(str)) == \"C1'\"\n        c1    = atom_array[mask]\n        if len(c1) == 0: return None\n        idx   = np.argsort(c1.res_id)\n        return c1[idx].coord\n    except Exception as e:\n        print(f\"    Error extracting C1' from {cif_file_path}: {e}\"); return None\n\ndef create_rnapro_input_json(sequence, target_id):\n    return [{\"sequences\": [{\"rnaSequence\": {\"sequence\": sequence, \"count\": 1}}], \"name\": target_id}]\n\ndef run_rnapro_single(target_id, sequence, configs, solution, template_idx, runner):\n    temp_dir = f\"./{configs.dump_dir}/input\"\n    os.makedirs(temp_dir, exist_ok=True)\n    json_path = os.path.join(temp_dir, f\"{target_id}_input.json\")\n    with open(json_path, \"w\") as f: json.dump(create_rnapro_input_json(sequence, target_id), f)\n    configs.input_json_path = json_path; configs.template_idx = int(template_idx)\n    infer_predict(runner, configs)\n    seed = configs.seeds[0] if hasattr(configs.seeds, '__iter__') else configs.seeds\n    cif_path = f\"{configs.dump_dir}/{target_id}/seed_{seed}/predictions/{target_id}_sample_0.cif\"\n    coord = extract_c1_coordinates(cif_path)\n    if coord is None: coord = np.zeros((len(sequence), 3), np.float32)\n    elif coord.shape[0] < len(sequence):\n        coord = np.concatenate([coord, np.zeros((len(sequence)-coord.shape[0], 3), np.float32)])\n    solution[target_id].coord.append(coord)\n\ndef run():\n    import random\n    from configs.configs_base import configs as configs_base_rnapro\n    from configs.configs_data import data_configs as data_configs_rnapro\n    from configs.configs_inference import inference_configs as inference_configs_rnapro\n    from rnapro.config import parse_sys_args\n\n    _seed = int(os.environ.get(\"RNAPRO_SEED\", \"62\"))\n    random.seed(_seed)\n    np.random.seed(_seed)\n    torch.manual_seed(_seed)\n    torch.cuda.manual_seed(_seed)\n    torch.cuda.manual_seed_all(_seed)\n    torch.backends.cudnn.benchmark = False\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.enabled = True\n    try:\n        torch.use_deterministic_algorithms(True)\n    except Exception as e:\n        print(f\"  [RNAPro] Warning: torch.use_deterministic_algorithms(True) failed: {e}\")\n        print(f\"  [RNAPro] Falling back to warn_only=True\")\n        try:\n            torch.use_deterministic_algorithms(True, warn_only=True)\n        except Exception:\n            pass\n\n    configs_base_rnapro[\"use_deepspeed_evo_attention\"] = \\\n        os.environ.get(\"USE_DEEPSPEED_EVO_ATTENTION\", False) == \"true\"\n    configs = {**configs_base_rnapro, **{\"data\": data_configs_rnapro}, **inference_configs_rnapro}\n    configs = parse_configs(configs=configs, arg_str=parse_sys_args(), fill_required_with_null=True)\n    valid_df = pd.read_csv(configs.sequences_csv)\n    print(f\"\\n -> {len(valid_df)} sequence(s) to process\")\n    runner  = InferenceRunner(configs)\n    solution = make_dummy_solution(valid_df)\n    for idx, row in valid_df.iterrows():\n        print(f\"\\n -> {row.target_id} (len={len(row.sequence)})\")\n        if len(row.sequence) > configs.max_len:\n            print(f\"    Skipping — too long (> {configs.max_len})\")\n            for _ in range(5): solution[row.target_id].coord.append(np.zeros((len(row.sequence),3), np.float32))\n            continue\n        try:\n            for template_idx in range(5):\n                run_rnapro_single(row.target_id, row.sequence, configs, solution, template_idx, runner)\n        except Exception as e:\n            print(f\"    ERROR: {e}\")\n\n    # ── Multi-GPU aware output path ───────────────────────────────────────────\n    _out = os.environ.get(\"RNAPRO_OUTPUT_CSV\", \"./submission_rnapro.csv\")\n    submit_df = solution_to_submit_df(solution).fillna(0.0)\n    submit_df.to_csv(_out, index=False)\n    print(f\"\\n -> Saved {_out}\")\n\nif __name__ == \"__main__\":\n    run()\n\"\"\"\n\n    with open(\"/kaggle/working/RNAPro/runner/inference.py\", \"w\") as f:\n        f.write(rnapro_runner_code)\n    print(\"  RNAPro runner script written\")\n\n\ndef _write_rnapro_bash_script(gpu_id, sequences_csv, dump_dir, output_csv):\n    \"\"\"Write one bash launcher for a single GPU.\"\"\"\n\n    script_name = f\"rnapro_inference_gpu{gpu_id}.sh\"\n    rnapro_sh = f\"\"\"#!/usr/bin/env bash\n\n# ── Pin to a single GPU ──────────────────────────────────────────────────────\nexport CUDA_VISIBLE_DEVICES={gpu_id}\n\n# ── Determinism environment variables ─────────────────────────────────────────\nexport LAYERNORM_TYPE=torch\nexport PYTHONHASHSEED={SEED}\nexport CUBLAS_WORKSPACE_CONFIG=\":4096:8\"\nexport RNAPRO_SEED={SEED}\nexport NVIDIA_TF32_OVERRIDE=0\nexport RNAPRO_OUTPUT_CSV=\"{output_csv}\"\n\nSEED={SEED}\nN_SAMPLE=1\nN_STEP=200\nN_CYCLE=10\n\nDUMP_DIR=\"{dump_dir}\"\nCHECKPOINT_PATH=\"../rnapro-private-best-500m.ckpt\"\nTEMPLATE_DATA=\"./release_data/kaggle/templates.pt\"\nRNA_MSA_DIR=\"/kaggle/input/stanford-rna-3d-folding-2/MSA\"\nSEQUENCES_CSV=\"{sequences_csv}\"\nRIBONANZA_PATH=\"/kaggle/input/models/shujun717/ribonanzanet2/pytorch/alpha/1/\"\nMODEL_NAME=\"{RNAPRO_MODEL_NAME}\"\n\nmkdir -p \"${{DUMP_DIR}}\"\n\npython3 runner/inference.py \\\\\n    --model_name \"${{MODEL_NAME}}\" \\\\\n    --seeds ${{SEED}} \\\\\n    --dump_dir \"${{DUMP_DIR}}\" \\\\\n    --load_checkpoint_path \"${{CHECKPOINT_PATH}}\" \\\\\n    --use_msa true \\\\\n    --use_template \"ca_precomputed\" \\\\\n    --model.use_template \"ca_precomputed\" \\\\\n    --model.use_RibonanzaNet2 true \\\\\n    --model.template_embedder.n_blocks 2 \\\\\n    --model.ribonanza_net_path \"${{RIBONANZA_PATH}}\" \\\\\n    --template_data \"${{TEMPLATE_DATA}}\" \\\\\n    --template_idx 0 \\\\\n    --rna_msa_dir \"${{RNA_MSA_DIR}}\" \\\\\n    --model.N_cycle ${{N_CYCLE}} \\\\\n    --sample_diffusion.N_sample ${{N_SAMPLE}} \\\\\n    --sample_diffusion.N_step ${{N_STEP}} \\\\\n    --load_strict true \\\\\n    --num_workers 0 \\\\\n    --triangle_attention \"torch\" \\\\\n    --triangle_multiplicative \"torch\" \\\\\n    --sequences_csv \"${{SEQUENCES_CSV}}\" \\\\\n    --deterministic true \\\\\n    --max_len {RNAPRO_MAX_LEN}\n\"\"\"\n    path = f\"/kaggle/working/RNAPro/{script_name}\"\n    with open(path, \"w\") as f:\n        f.write(rnapro_sh)\n    return path\n\n\n\n# ═══════════════════════════════════════════════════════════════════════════════\n# §13  Phase 3 — RNAPro Refinement on ALL 5 combined predictions\n#      (Copied from Script 1, runs AFTER Protenix+Boltz2 ensemble)\n# ═══════════════════════════════════════════════════════════════════════════════\n\ndef _export_combined_as_rnapro_templates(test_df, combined_preds):\n    \"\"\"Export all N_SAMPLE combined predictions to CSV → .pt for RNAPro.\"\"\"\n    rows = []\n    for _, row in test_df.iterrows():\n        tid = row[\"target_id\"]\n        seq = row[\"sequence\"]\n        preds = combined_preds.get(tid, [])\n        for j in range(len(seq)):\n            r = {\"ID\": f\"{tid}_{j+1}\", \"resname\": seq[j], \"resid\": j + 1}\n            for i in range(N_SAMPLE):\n                if i < len(preds) and j < preds[i].shape[0]:\n                    r[f\"x_{i+1}\"] = float(preds[i][j][0])\n                    r[f\"y_{i+1}\"] = float(preds[i][j][1])\n                    r[f\"z_{i+1}\"] = float(preds[i][j][2])\n                else:\n                    r[f\"x_{i+1}\"] = 0.0\n                    r[f\"y_{i+1}\"] = 0.0\n                    r[f\"z_{i+1}\"] = 0.0\n            rows.append(r)\n\n    sub_df = pd.DataFrame(rows)\n    col_order = [\"ID\", \"resname\", \"resid\"] + [\n        f\"{c}_{i}\" for i in range(1, N_SAMPLE + 1) for c in [\"x\", \"y\", \"z\"]\n    ]\n    sub_df = sub_df[col_order]\n    csv_path = \"/kaggle/working/submission_combined_for_rnapro.csv\"\n    sub_df.to_csv(csv_path, index=False)\n    print(f\"  Combined predictions exported to {csv_path}\")\n\n    os.chdir(\"/kaggle/working/RNAPro\")\n    os.system(\n        f\"python preprocess/convert_templates_to_pt_files.py \"\n        f\"--input_csv {csv_path} --output_name templates.pt\"\n    )\n    os.chdir(\"/kaggle/working\")\n    print(\"  Combined templates converted to .pt for RNAPro\")\n\n\ndef _parse_rnapro_results_combined(test_df, combined_preds, rnapro_csv_paths) -> int:\n    \"\"\"\n    Parse RNAPro CSV outputs and replace combined predictions in-place.\n    Returns count of targets refined.\n    \"\"\"\n    frames = []\n    for csv_path in rnapro_csv_paths:\n        if os.path.isfile(csv_path):\n            frames.append(pd.read_csv(csv_path))\n        else:\n            print(f\"  ⚠ {csv_path} not found — skipping\")\n\n    if not frames:\n        print(\"  ⚠ No RNAPro output CSVs found — keeping all combined predictions\")\n        return 0\n\n    df_rnapro = pd.concat(frames, ignore_index=True)\n    df_rnapro[\"target_id\"] = df_rnapro[\"ID\"].str.rsplit(\"_\", n=1).str[0]\n\n    n_refined = 0\n    for tid, grp in df_rnapro.groupby(\"target_id\", sort=False):\n        grp = grp.sort_values(\"resid\")\n        seq_len = len(test_df[test_df[\"target_id\"] == tid][\"sequence\"].values[0])\n\n        refined = []\n        for k in range(1, N_SAMPLE + 1):\n            cols = [f\"x_{k}\", f\"y_{k}\", f\"z_{k}\"]\n            if not all(c in grp.columns for c in cols):\n                continue\n            arr = grp[cols].values.astype(np.float64)\n            if (arr.shape[0] == seq_len\n                    and np.isfinite(arr).all()\n                    and not np.all(arr == 0)):\n                refined.append(arr)\n\n        if not refined:\n            print(f\"  {tid}: RNAPro produced no valid outputs, keeping combined preds\")\n            continue\n\n        old_preds = combined_preds.get(tid, [])\n        new_preds = []\n        for i in range(N_SAMPLE):\n            if i < len(refined):\n                new_preds.append(refined[i])\n            elif i < len(old_preds):\n                new_preds.append(old_preds[i])\n            else:\n                new_preds.append(np.zeros((seq_len, 3), dtype=np.float64))\n\n        combined_preds[tid] = new_preds[:N_SAMPLE]\n        n_refined += 1\n        n_used = min(len(refined), N_SAMPLE)\n        print(f\"  {tid}: {n_used}/{N_SAMPLE} predictions refined by RNAPro\")\n\n    return n_refined\n\n\ndef rnapro_refine_all(test_df, combined_preds):\n    \"\"\"\n    Phase 3 — Run RNAPro to refine ALL combined predictions in-place.\n    Parallelizes across available GPUs.\n    Modifies combined_preds dict directly.\n    Returns number of targets refined.\n    \"\"\"\n    n_gpus = _get_n_gpus()\n\n    print(f\"\\n{'=' * 60}\")\n    print(f\"PHASE 3: RNAPro Refinement of ALL Combined Predictions\")\n    print(f\"  RNAPRO_MAX_LEN = {RNAPRO_MAX_LEN}  |  N_GPUS = {n_gpus}\")\n    print(f\"{'=' * 60}\")\n    t0 = time.time()\n\n    eligible = test_df[test_df[\"sequence\"].str.len() <= RNAPRO_MAX_LEN]\n    skipped = test_df[test_df[\"sequence\"].str.len() > RNAPRO_MAX_LEN]\n\n    if len(skipped) > 0:\n        for _, row in skipped.iterrows():\n            print(f\"  {row['target_id']} ({len(row['sequence'])} nt): \"\n                  f\"too long for RNAPro, keeping combined predictions as-is\")\n\n    if len(eligible) == 0:\n        print(\"  No targets eligible for RNAPro — skipping\")\n        return 0\n\n    # Step 1: Export all 5 combined predictions as RNAPro templates\n    print(f\"\\n  Exporting {len(eligible)} targets' combined predictions as RNAPro templates...\")\n    _export_combined_as_rnapro_templates(test_df, combined_preds)\n\n    # Step 2: Write runner script (shared)\n    _write_rnapro_runner_code()\n\n    # Step 3: Split targets across GPUs\n    effective_gpus = min(n_gpus, len(eligible))\n    splits = [eligible.iloc[i::effective_gpus].reset_index(drop=True)\n              for i in range(effective_gpus)]\n\n    # Step 4: Write per-GPU sequences CSVs and bash scripts\n    script_paths = []\n    output_csv_paths = []\n    for gpu_id in range(effective_gpus):\n        seq_csv = f\"/kaggle/working/rnapro_sequences_gpu{gpu_id}.csv\"\n        dump_dir = f\"../rnapro_output_gpu{gpu_id}\"\n        out_csv = f\"/kaggle/working/rnapro_output_gpu{gpu_id}/submission_rnapro.csv\"\n\n        splits[gpu_id].to_csv(seq_csv, index=False)\n        script_path = _write_rnapro_bash_script(gpu_id, seq_csv, dump_dir, out_csv)\n        script_paths.append(script_path)\n        output_csv_paths.append(out_csv)\n\n        n_tgt = len(splits[gpu_id])\n        print(f\"  GPU {gpu_id}: {n_tgt} target(s) → {seq_csv}\")\n\n    # Step 5: Launch all GPU workers in parallel\n    print(f\"\\n  Launching {effective_gpus} RNAPro worker(s) in parallel...\")\n    os.chdir(\"/kaggle/working/RNAPro\")\n\n    procs = []\n    for gpu_id in range(effective_gpus):\n        p = _subprocess.Popen(\n            [\"bash\", script_paths[gpu_id]],\n            stdout=_subprocess.PIPE,\n            stderr=_subprocess.STDOUT,\n        )\n        procs.append(p)\n        print(f\"  GPU {gpu_id}: PID {p.pid} started\")\n\n    for gpu_id, p in enumerate(procs):\n        stdout, _ = p.communicate()\n        if stdout:\n            for line in stdout.decode(\"utf-8\", errors=\"replace\").splitlines():\n                print(f\"  [GPU{gpu_id}] {line}\")\n        if p.returncode != 0:\n            print(f\"  ⚠ GPU {gpu_id} RNAPro returned exit code {p.returncode}\")\n\n    os.chdir(\"/kaggle/working\")\n\n    # Step 6: Parse merged results\n    n_refined = _parse_rnapro_results_combined(test_df, combined_preds, output_csv_paths)\n\n    # Step 7: Free GPU memory\n    gc.collect()\n    torch.cuda.empty_cache()\n    gc.collect()\n\n    elapsed = time.time() - t0\n    print(f\"\\nPhase 3 done in {elapsed:.1f}s  ({effective_gpus} GPU(s))\")\n    print(f\"  Refined: {n_refined}/{len(eligible)} eligible targets\")\n    return n_refined","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ═══════════════════════════════════════════════════════════════════════════════\n# §12  Phase 2 — Protenix (Two-phase: ensemble heavy, then split light)\n# ═══════════════════════════════════════════════════════════════════════════════\n#\n# Phase A: targets needing >= ENSEMBLE_THRESHOLD slots\n#   GPU 0 → MSA,  GPU 1 → noMSA  (same targets)\n#   Wait → Kabsch-align + average\n#\n# Phase B: targets needing < ENSEMBLE_THRESHOLD slots\n#   GPU 0 → half (MSA),  GPU 1 → other half (MSA)\n#   Wait → collect\n#\n# Both GPUs stay busy in both phases.  Split targets get proper MSA.\n# ═══════════════════════════════════════════════════════════════════════════════\n\nENSEMBLE_THRESHOLD = 3\n\nimport copy as _copy\nimport pickle as _pickle\nimport subprocess as _subprocess\n\n\n# ── Chunking helpers ──────────────────────────────────────────────────────────\n\ndef chunk_sequence(seq_len: int, max_len: int, overlap: int) -> list:\n    if seq_len <= max_len:\n        return [(0, seq_len)]\n    chunks, start = [], 0\n    while start < seq_len:\n        end = min(start + max_len, seq_len)\n        chunks.append((start, end))\n        if end == seq_len:\n            break\n        start = end - overlap\n        if seq_len - start < overlap + 20:\n            chunks[-1] = (chunks[-1][0], seq_len)\n            break\n    return chunks\n\n\ndef kabsch_align(mobile, target):\n    assert mobile.shape == target.shape and mobile.shape[1] == 3\n    cm, ct = mobile.mean(0), target.mean(0)\n    H = (mobile - cm).T @ (target - ct)\n    U, S, Vt = np.linalg.svd(H)\n    d = np.linalg.det(Vt.T @ U.T)\n    R = Vt.T @ np.diag([1, 1, np.sign(d)]) @ U.T\n    return R, ct - cm @ R.T\n\n\ndef apply_transform(coords, R, t):\n    return coords @ R.T + t\n\n\ndef blend_weights(n):\n    return np.linspace(1.0, 0.0, n)[:, None]\n\n\ndef assemble_chunks_single_sample(chunk_coords, chunk_ranges, full_len, overlap):\n    if len(chunk_coords) == 1:\n        out = np.zeros((full_len, 3), dtype=np.float32)\n        s, e = chunk_ranges[0]\n        out[s:e] = chunk_coords[0][:e - s]\n        return out\n    assembled = np.full((full_len, 3), np.nan, dtype=np.float32)\n    s0, e0 = chunk_ranges[0]\n    assembled[s0:e0] = chunk_coords[0][:e0 - s0]\n    for i in range(1, len(chunk_coords)):\n        s_p, e_p = chunk_ranges[i - 1]\n        s_c, e_c = chunk_ranges[i]\n        cc = chunk_coords[i][:e_c - s_c].copy()\n        ov_s, ov_e = s_c, min(e_p, e_c)\n        ov_len = ov_e - ov_s\n        if ov_len < 3:\n            assembled[ov_e:e_c] = cc[ov_e - s_c:]\n            continue\n        tgt = assembled[ov_s:ov_e].copy()\n        mob = cc[:ov_len].copy()\n        v = (np.isfinite(tgt).all(1) & np.isfinite(mob).all(1)\n             & (np.abs(tgt).sum(1) > 1e-6) & (np.abs(mob).sum(1) > 1e-6))\n        if v.sum() < 3:\n            assembled[ov_e:e_c] = cc[ov_len:]\n            continue\n        R, t = kabsch_align(mob[v], tgt[v])\n        ca = apply_transform(cc, R, t)\n        w = blend_weights(ov_len)\n        assembled[ov_s:ov_e] = w * assembled[ov_s:ov_e] + (1 - w) * ca[:ov_len]\n        if ov_len < (e_c - s_c):\n            assembled[ov_e:e_c] = ca[ov_len:]\n    return np.nan_to_num(assembled, nan=0.0)\n\n\n# ── Worker script ─────────────────────────────────────────────────────────────\n\ndef _write_protenix_worker_code():\n    worker_code = r\"\"\"\nimport os, gc, json, pickle, random, argparse, hashlib\nfrom pathlib import Path\nimport numpy as np\nimport torch\n\n\ndef seed_everything_local(seed):\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    os.environ[\"CUBLAS_WORKSPACE_CONFIG\"] = \":4096:8\"\n    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed); torch.cuda.manual_seed_all(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\n\ndef build_configs(input_json_path, dump_dir, model_name, seed,\n                  use_msa, use_template, use_rna_msa, model_n_sample):\n    from configs.configs_base import configs as cb\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    base = {**cb, **{\"data\": data_configs}, **inference_configs}\n    def _du(t, p):\n        for k, v in p.items():\n            if isinstance(v, dict) and k in t and isinstance(t[k], dict): _du(t[k], v)\n            else: t[k] = v\n    _du(base, model_configs[model_name])\n    arg_str = (f\"--model_name {model_name} --input_json_path {input_json_path} \"\n               f\"--dump_dir {dump_dir} --use_msa {use_msa} --use_template {use_template} \"\n               f\"--use_rna_msa {use_rna_msa} --sample_diffusion.N_sample {model_n_sample} \"\n               f\"--seeds {seed}\")\n    return parse_configs(configs=base, arg_str=arg_str, fill_required_with_null=True)\n\n\ndef chunk_sequence(seq_len, max_len, overlap):\n    if seq_len <= max_len: return [(0, seq_len)]\n    chunks, start = [], 0\n    while start < seq_len:\n        end = min(start + max_len, seq_len)\n        chunks.append((start, end))\n        if end == seq_len: break\n        start = end - overlap\n        if seq_len - start < overlap + 20:\n            chunks[-1] = (chunks[-1][0], seq_len); break\n    return chunks\n\n\ndef kabsch_align(mobile, target):\n    cm, ct = mobile.mean(0), target.mean(0)\n    H = (mobile - cm).T @ (target - ct)\n    U, S, Vt = np.linalg.svd(H)\n    d = np.linalg.det(Vt.T @ U.T)\n    R = Vt.T @ np.diag([1, 1, np.sign(d)]) @ U.T\n    return R, ct - cm @ R.T\n\ndef apply_transform(c, R, t): return c @ R.T + t\ndef blend_weights(n): return np.linspace(1.0, 0.0, n)[:, None]\n\ndef assemble_chunks_single_sample(chunk_coords, chunk_ranges, full_len, overlap):\n    if len(chunk_coords) == 1:\n        out = np.zeros((full_len, 3), np.float32)\n        s, e = chunk_ranges[0]; out[s:e] = chunk_coords[0][:e-s]; return out\n    asm = np.full((full_len, 3), np.nan, np.float32)\n    s0, e0 = chunk_ranges[0]; asm[s0:e0] = chunk_coords[0][:e0-s0]\n    for i in range(1, len(chunk_coords)):\n        sp, ep = chunk_ranges[i-1]; sc, ec = chunk_ranges[i]\n        cc = chunk_coords[i][:ec-sc].copy()\n        os_, oe = sc, min(ep, ec); ol = oe - os_\n        if ol < 3: asm[oe:ec] = cc[ol:]; continue\n        tgt, mob = asm[os_:oe].copy(), cc[:ol].copy()\n        v = (np.isfinite(tgt).all(1)&np.isfinite(mob).all(1)\n             &(np.abs(tgt).sum(1)>1e-6)&(np.abs(mob).sum(1)>1e-6))\n        if v.sum()<3: asm[oe:ec] = cc[ol:]; continue\n        R,t = kabsch_align(mob[v],tgt[v]); ca = apply_transform(cc,R,t)\n        w = blend_weights(ol)\n        asm[os_:oe] = w*asm[os_:oe]+(1-w)*ca[:ol]\n        if ol<(ec-sc): asm[oe:ec] = ca[ol:]\n    return np.nan_to_num(asm, nan=0.0)\n\n\ndef _resolve_c1_mask(feat, seq, raw):\n    if \"centre_atom_mask\" in feat:\n        return (feat[\"centre_atom_mask\"]==1).to(raw.device)\n    if \"atom_to_tokatom_idx\" in feat:\n        m11=(feat[\"atom_to_tokatom_idx\"]==11).to(raw.device)\n        m12=(feat[\"atom_to_tokatom_idx\"]==12).to(raw.device)\n        return m11 if abs(m11.sum()-len(seq))<abs(m12.sum()-len(seq)) else m12\n    return torch.zeros(raw.shape[1],dtype=torch.bool,device=raw.device)\n\n\ndef _predict_chunk(runner, configs, update_fn, chunk_seq, n_samples, tid, ci, DS):\n    work = Path(\"/kaggle/working\")\n    name = f\"{tid}_chunk{ci}\"\n    jpath = str(work/f\"ptx_chunk_{name}_{os.getpid()}.json\")\n    with open(jpath,\"w\") as f:\n        json.dump([{\"name\":name,\"covalent_bonds\":[],\n                     \"sequences\":[{\"rnaSequence\":{\"sequence\":chunk_seq,\"count\":1}}]}],f)\n    old = configs.input_json_path; configs.input_json_path = jpath\n    try:\n        ds = DS(configs)\n        if len(ds)==0: return None\n        data, aa, err = ds[0]\n        if err: print(f\"  [CHUNK] {name}: {err}\"); return None\n        cfg2 = update_fn(configs, data[\"N_token\"].item())\n        cfg2.sample_diffusion.N_sample = n_samples\n        runner.update_model_configs(cfg2)\n        pred = runner.predict(data)\n        raw = pred[\"coordinate\"]; feat = data[\"input_feature_dict\"]\n        mask = _resolve_c1_mask(feat, chunk_seq, raw)\n        coords = raw[:,mask,:].detach().cpu().numpy()\n        if coords.shape[1]!=len(chunk_seq):\n            p = np.zeros((coords.shape[0],len(chunk_seq),3),np.float32)\n            mn = min(coords.shape[1],len(chunk_seq))\n            p[:,:mn,:] = coords[:,:mn,:]\n            coords = p\n        return coords\n    except Exception as e:\n        import traceback; print(f\"  [CHUNK] {name}: FAILED — {e}\"); traceback.print_exc()\n        return None\n    finally:\n        configs.input_json_path = old\n        gc.collect(); torch.cuda.empty_cache()\n        try: os.remove(jpath)\n        except: pass\n\n\ndef _predict_full(runner, configs, update_fn, seq, tid, chunks, n_samples, DS):\n    L = len(seq); nc = len(chunks)\n    if nc == 1:\n        cs, ce = chunks[0]\n        coords = _predict_chunk(runner,configs,update_fn,seq[cs:ce],n_samples,tid,0,DS)\n        if coords is None: return None\n        if ce-cs < L:\n            p = np.zeros((coords.shape[0],L,3),np.float32)\n            p[:,cs:ce,:] = coords[:,:ce-cs,:]; return p\n        return coords\n    all_cc = []\n    for ci,(cs,ce) in enumerate(chunks):\n        cc = _predict_chunk(runner,configs,update_fn,seq[cs:ce],n_samples,tid,ci,DS)\n        if cc is None or cc.shape[0]<n_samples: return None\n        all_cc.append(cc)\n    out = np.zeros((n_samples,L,3),np.float32)\n    for si in range(n_samples):\n        out[si] = assemble_chunks_single_sample(\n            [all_cc[ci][si] for ci in range(nc)], chunks, L, 64)\n    return out\n\n\ndef main():\n    ap = argparse.ArgumentParser()\n    ap.add_argument(\"--tasks_pkl\",required=True)\n    ap.add_argument(\"--result_pkl\",required=True)\n    ap.add_argument(\"--code_dir\",required=True)\n    ap.add_argument(\"--root_dir\",required=True)\n    ap.add_argument(\"--model_name\",required=True)\n    ap.add_argument(\"--seed\",type=int,required=True)\n    ap.add_argument(\"--max_seq_len\",type=int,required=True)\n    ap.add_argument(\"--chunk_overlap\",type=int,required=True)\n    ap.add_argument(\"--n_sample\",type=int,required=True)\n    ap.add_argument(\"--use_msa\",required=True)\n    ap.add_argument(\"--use_template\",required=True)\n    ap.add_argument(\"--use_rna_msa\",required=True)\n    ap.add_argument(\"--model_n_sample\",type=int,required=True)\n    args = ap.parse_args()\n\n    os.environ[\"PROTENIX_ROOT_DIR\"] = args.root_dir\n    import sys\n    sys.path = [p for p in sys.path if \"/kaggle/working/RNAPro\" not in str(p)]\n    if args.code_dir in sys.path: sys.path.remove(args.code_dir)\n    sys.path.insert(0, args.code_dir)\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    with open(args.tasks_pkl,\"rb\") as f: tasks = pickle.load(f)\n\n    if not tasks:\n        with open(args.result_pkl,\"wb\") as f: pickle.dump({},f)\n        return\n\n    work = Path(\"/kaggle/working\")\n    init_json = str(work/f\"ptx_init_{os.getpid()}.json\")\n    dummy = tasks[0][\"sequence\"][:args.max_seq_len]\n    with open(init_json,\"w\") as f:\n        json.dump([{\"name\":\"init\",\"covalent_bonds\":[],\n                     \"sequences\":[{\"rnaSequence\":{\"sequence\":dummy,\"count\":1}}]}],f)\n\n    seed_everything_local(args.seed)\n    cfg = build_configs(init_json, str(work/f\"out_{os.getpid()}\"), args.model_name,\n                        args.seed, args.use_msa, args.use_template, args.use_rna_msa,\n                        args.model_n_sample)\n    cfg = update_gpu_compatible_configs(cfg)\n    runner = InferenceRunner(cfg)\n\n    results = {}\n    for task in tasks:\n        tid = task[\"target_id\"]; seq = task[\"sequence\"]; n_needed = int(task[\"n_needed\"])\n        chunks = chunk_sequence(len(seq), args.max_seq_len, args.chunk_overlap)\n        print(f\"\\n[{os.getpid()}] {tid} ({len(seq)} nt): {n_needed} samples, \"\n              f\"{'single chunk' if len(chunks)==1 else f'{len(chunks)} chunks'}\")\n\n        task_seed = args.seed + (int(hashlib.sha256(tid.encode()).hexdigest(),16)%(2**63)) % 100000\n        seed_everything_local(task_seed)\n\n        coords = _predict_full(runner, cfg, update_inference_configs,\n                               seq, tid, chunks, max(n_needed, args.n_sample),\n                               InferenceDataset)\n\n        if coords is not None:\n            results[tid] = {\"samples\": coords[:n_needed]}\n            print(f\"[{os.getpid()}] {tid}: OK shape={coords[:n_needed].shape}\")\n        else:\n            results[tid] = {\"samples\": None}\n            print(f\"[{os.getpid()}] {tid}: FAILED\")\n\n        gc.collect(); torch.cuda.empty_cache()\n\n    with open(args.result_pkl,\"wb\") as f: pickle.dump(results,f)\n    try: os.remove(init_json)\n    except: pass\n\nif __name__==\"__main__\": main()\n\"\"\"\n    path = \"/kaggle/working/protenix_worker.py\"\n    with open(path, \"w\", encoding=\"utf-8\") as f:\n        f.write(worker_code)\n    return path\n\n\n# ── Helpers ───────────────────────────────────────────────────────────────────\n\ndef _parse_protenix_worker_results(result_paths):\n    merged = {}\n    for p in result_paths:\n        if not os.path.isfile(p):\n            print(f\"  ⚠ Missing result: {p}\"); continue\n        with open(p, \"rb\") as f:\n            merged.update(_pickle.load(f))\n    return merged\n\n\ndef _ensemble_coords(coords_msa, coords_nomsa, weight_msa=0.5):\n    aligned = _kabsch_superpose(coords_nomsa.copy(), coords_msa.copy())\n    return weight_msa * coords_msa + (1.0 - weight_msa) * aligned\n\n\ndef _match_and_ensemble(msa_samples, nomsa_samples, n_select, weight_msa=0.5):\n    M, N = msa_samples.shape[0], nomsa_samples.shape[0]\n    tm_mat = np.zeros((M, N), dtype=np.float64)\n    for i in range(M):\n        for j in range(N):\n            tm_mat[i, j] = tm_score_rna(msa_samples[i], nomsa_samples[j])\n    used_m, used_n, pairs = set(), set(), []\n    for flat_idx in np.argsort(tm_mat.ravel())[::-1]:\n        mi, ni = int(flat_idx // N), int(flat_idx % N)\n        if mi in used_m or ni in used_n: continue\n        pairs.append((mi, ni)); used_m.add(mi); used_n.add(ni)\n        if len(pairs) >= n_select: break\n    for mi in range(M):\n        if mi not in used_m and len(pairs) < n_select:\n            pairs.append((mi, None))\n    results = []\n    for mi, ni in pairs[:n_select]:\n        if ni is not None:\n            results.append(_ensemble_coords(msa_samples[mi], nomsa_samples[ni], weight_msa))\n        else:\n            results.append(msa_samples[mi].copy())\n    return np.stack(results)\n\n\ndef _make_tasks(targets_dict, template_preds):\n    tasks = []\n    for tid, (n_needed, seq) in targets_dict.items():\n        tbm = template_preds.get(tid, []) if template_preds else []\n        tasks.append({\n            \"target_id\": tid, \"sequence\": seq, \"n_needed\": int(n_needed),\n            \"tbm_coords\": [np.asarray(x, np.float32) for x in tbm],\n        })\n    return tasks\n\n\ndef _launch_worker(gpu_id, label, task_pkl, result_pkl, worker_py,\n                   use_msa, use_rna_msa):\n    env = os.environ.copy()\n    env[\"CUDA_VISIBLE_DEVICES\"] = str(gpu_id)\n    env[\"PYTHONHASHSEED\"] = str(SEED)\n    env[\"CUBLAS_WORKSPACE_CONFIG\"] = \":4096:8\"\n    cmd = [\n        sys.executable, worker_py,\n        \"--tasks_pkl\", task_pkl, \"--result_pkl\", result_pkl,\n        \"--code_dir\", DEFAULT_CODE_DIR, \"--root_dir\", DEFAULT_ROOT_DIR,\n        \"--model_name\", MODEL_NAME, \"--seed\", str(SEED),\n        \"--max_seq_len\", str(MAX_SEQ_LEN),\n        \"--chunk_overlap\", str(CHUNK_OVERLAP),\n        \"--n_sample\", str(N_SAMPLE),\n        \"--use_msa\", use_msa,\n        \"--use_template\", USE_TEMPLATE,\n        \"--use_rna_msa\", use_rna_msa,\n        \"--model_n_sample\", str(MODEL_N_SAMPLE),\n    ]\n    p = _subprocess.Popen(cmd, stdout=_subprocess.PIPE, stderr=_subprocess.STDOUT,\n                          env=env, text=True, cwd=DEFAULT_CODE_DIR)\n    print(f\"    GPU {gpu_id} ({label}): PID {p.pid}\")\n    return p\n\n\ndef _wait_and_collect(procs, result_paths):\n    for gid, label, p in procs:\n        stdout, _ = p.communicate()\n        if stdout:\n            for line in stdout.splitlines():\n                print(f\"    [GPU{gid}/{label}] {line}\")\n        if p.returncode != 0:\n            print(f\"    ⚠ GPU {gid} ({label}) exit code {p.returncode}\")\n    merged = {}\n    for rp in result_paths:\n        merged.update(_parse_protenix_worker_results([rp]))\n    return merged\n\n\n# ── Main phase ────────────────────────────────────────────────────────────────\n\ndef protenix_phase(protenix_queue, test_df_trunc, configs_builder_fn,\n                   template_preds=None):\n    if not protenix_queue:\n        return {}\n\n    n_gpus = _get_n_gpus()\n    assert n_gpus >= 2, f\"Need 2 GPUs, have {n_gpus}\"\n\n    # ── Partition ─────────────────────────────────────────────────────────────\n    ensemble_targets = {}\n    split_targets = {}\n    for tid, (n_needed, seq) in protenix_queue.items():\n        if n_needed >= ENSEMBLE_THRESHOLD:\n            ensemble_targets[tid] = (n_needed, seq)\n        else:\n            split_targets[tid] = (n_needed, seq)\n\n    print(f\"\\n{'=' * 60}\")\n    print(f\"PHASE 2: Two-Phase Protenix for {len(protenix_queue)} targets\")\n    print(f\"  Phase A — ensemble (n≥{ENSEMBLE_THRESHOLD}): {len(ensemble_targets)} targets\")\n    print(f\"  Phase B — split   (n<{ENSEMBLE_THRESHOLD}):  {len(split_targets)} targets\")\n    print(f\"  seed={SEED}  MAX_SEQ_LEN={MAX_SEQ_LEN}\")\n    print(f\"{'=' * 60}\")\n\n    worker_py = _write_protenix_worker_code()\n    protenix_preds = {}\n    tmp_files = [worker_py]\n\n    # ══════════════════════════════════════════════════════════════════════════\n    # PHASE A: Ensemble targets — GPU 0 MSA, GPU 1 noMSA (same targets)\n    # ══════════════════════════════════════════════════════════════════════════\n    if ensemble_targets:\n        print(f\"\\n  ── Phase A: Ensemble {len(ensemble_targets)} targets ──\")\n        ens_tasks = _make_tasks(ensemble_targets, template_preds)\n\n        tp_a0 = \"/kaggle/working/ptx_ens_msa.pkl\"\n        tp_a1 = \"/kaggle/working/ptx_ens_nomsa.pkl\"\n        rp_a0 = \"/kaggle/working/ptx_res_ens_msa.pkl\"\n        rp_a1 = \"/kaggle/working/ptx_res_ens_nomsa.pkl\"\n\n        with open(tp_a0, \"wb\") as f: _pickle.dump(ens_tasks, f)\n        with open(tp_a1, \"wb\") as f: _pickle.dump(ens_tasks, f)\n        tmp_files += [tp_a0, tp_a1, rp_a0, rp_a1]\n\n        p0 = _launch_worker(0, \"ENS/MSA\",   tp_a0, rp_a0, worker_py, \"true\",  \"true\")\n        p1 = _launch_worker(1, \"ENS/noMSA\", tp_a1, rp_a1, worker_py, \"false\", \"false\")\n\n        res_a = {}\n        for gid, label, p in [(0, \"ENS/MSA\", p0), (1, \"ENS/noMSA\", p1)]:\n            stdout, _ = p.communicate()\n            if stdout:\n                for line in stdout.splitlines():\n                    print(f\"    [GPU{gid}/{label}] {line}\")\n            if p.returncode != 0:\n                print(f\"    ⚠ GPU {gid} ({label}) exit code {p.returncode}\")\n\n        res_msa = _parse_protenix_worker_results([rp_a0])\n        res_nomsa = _parse_protenix_worker_results([rp_a1])\n\n        print(f\"    MSA:   {sum(1 for v in res_msa.values()   if v.get('samples') is not None)}/{len(ensemble_targets)} ok\")\n        print(f\"    noMSA: {sum(1 for v in res_nomsa.values() if v.get('samples') is not None)}/{len(ensemble_targets)} ok\")\n\n        # Ensemble\n        for tid, (n_needed, seq) in ensemble_targets.items():\n            msa_arr = res_msa.get(tid, {}).get(\"samples\")\n            nomsa_arr = res_nomsa.get(tid, {}).get(\"samples\")\n            has_m = msa_arr is not None and msa_arr.ndim == 3\n            has_n = nomsa_arr is not None and nomsa_arr.ndim == 3\n\n            if has_m and has_n:\n                ens = _match_and_ensemble(msa_arr, nomsa_arr, n_needed, 0.5)\n                protenix_preds[tid] = ens\n                ptms = [tm_score_rna(msa_arr[i], nomsa_arr[i])\n                        for i in range(min(msa_arr.shape[0], nomsa_arr.shape[0], n_needed))]\n                print(f\"    {tid}: ensembled {ens.shape[0]} (avg pair TM={np.mean(ptms):.3f})\")\n            elif has_m:\n                protenix_preds[tid] = msa_arr[:n_needed]\n                print(f\"    {tid}: MSA only\")\n            elif has_n:\n                protenix_preds[tid] = nomsa_arr[:n_needed]\n                print(f\"    {tid}: noMSA only\")\n            else:\n                protenix_preds[tid] = None\n                print(f\"    {tid}: FAILED both\")\n\n        # Free GPU memory between phases\n        gc.collect(); torch.cuda.empty_cache()\n\n    # ══════════════════════════════════════════════════════════════════════════\n    # PHASE B: Split targets — round-robin across GPUs, both MSA\n    # ══════════════════════════════════════════════════════════════════════════\n    if split_targets:\n        print(f\"\\n  ── Phase B: Split {len(split_targets)} targets ──\")\n\n        # Round-robin by descending sequence length\n        split_items = sorted(split_targets.items(), key=lambda x: len(x[1][1]), reverse=True)\n        split_g0 = {tid: v for i, (tid, v) in enumerate(split_items) if i % 2 == 0}\n        split_g1 = {tid: v for i, (tid, v) in enumerate(split_items) if i % 2 == 1}\n\n        tasks_g0 = _make_tasks(split_g0, template_preds)\n        tasks_g1 = _make_tasks(split_g1, template_preds)\n\n        print(f\"    GPU 0: {len(tasks_g0)} targets (MSA)\")\n        print(f\"    GPU 1: {len(tasks_g1)} targets (MSA)\")\n\n        tp_b0 = \"/kaggle/working/ptx_split_g0.pkl\"\n        tp_b1 = \"/kaggle/working/ptx_split_g1.pkl\"\n        rp_b0 = \"/kaggle/working/ptx_res_split_g0.pkl\"\n        rp_b1 = \"/kaggle/working/ptx_res_split_g1.pkl\"\n        tmp_files += [tp_b0, tp_b1, rp_b0, rp_b1]\n\n        procs_b = []\n\n        if tasks_g0:\n            with open(tp_b0, \"wb\") as f: _pickle.dump(tasks_g0, f)\n            p0 = _launch_worker(0, \"SPLIT/MSA\", tp_b0, rp_b0, worker_py, \"true\", \"true\")\n            procs_b.append((0, \"SPLIT/MSA\", p0))\n        else:\n            # Write empty result\n            with open(rp_b0, \"wb\") as f: _pickle.dump({}, f)\n\n        if tasks_g1:\n            with open(tp_b1, \"wb\") as f: _pickle.dump(tasks_g1, f)\n            p1 = _launch_worker(1, \"SPLIT/MSA\", tp_b1, rp_b1, worker_py, \"true\", \"true\")\n            procs_b.append((1, \"SPLIT/MSA\", p1))\n        else:\n            with open(rp_b1, \"wb\") as f: _pickle.dump({}, f)\n\n        for gid, label, p in procs_b:\n            stdout, _ = p.communicate()\n            if stdout:\n                for line in stdout.splitlines():\n                    print(f\"    [GPU{gid}/{label}] {line}\")\n            if p.returncode != 0:\n                print(f\"    ⚠ GPU {gid} ({label}) exit code {p.returncode}\")\n\n        res_b = {}\n        res_b.update(_parse_protenix_worker_results([rp_b0]))\n        res_b.update(_parse_protenix_worker_results([rp_b1]))\n\n        for tid, (n_needed, seq) in split_targets.items():\n            arr = res_b.get(tid, {}).get(\"samples\")\n            if arr is not None and arr.ndim == 3:\n                protenix_preds[tid] = arr[:n_needed]\n                print(f\"    {tid}: {arr[:n_needed].shape[0]} samples\")\n            else:\n                protenix_preds[tid] = None\n                print(f\"    {tid}: FAILED\")\n\n    # ── Cleanup ───────────────────────────────────────────────────────────────\n    for p in tmp_files:\n        try: os.remove(p)\n        except: pass\n    gc.collect(); torch.cuda.empty_cache()\n\n    return protenix_preds","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ═══════════════════════════════════════════════════════════════════════════════\n# §12.5  Phase 2.5 — Boltz2  (NO CHUNKING — one predict call per slot)\n# ═══════════════════════════════════════════════════════════════════════════════\n\ndef _write_boltz_gpu_worker_code():\n    \"\"\"Write a persistent Boltz2 worker script.\n\n    ONE process per GPU.  Loads torch + boltz + model ONCE, then processes\n    every job assigned to that GPU sequentially via Click CliRunner.\n    No chunking — each job is one full-sequence predict call.\n    \"\"\"\n\n    worker_code = r'''#!/usr/bin/env python3\n\"\"\"\nPersistent Boltz2 GPU worker.\nLoads the model once, then runs all assigned jobs in-process via CliRunner.\nNo chunking — one predict call per (target, repeat).\n\"\"\"\nimport os, sys, gc, pickle, argparse, shutil, traceback\nimport numpy as np\nimport torch\n\n# ── Patch torch.load ONCE + cache checkpoint loads ───────────────────────────\n_real_torch_load = torch.load\n_ckpt_cache = {}\n\ndef _cached_torch_load(*args, **kwargs):\n    kwargs[\"weights_only\"] = False\n    path = args[0] if args else kwargs.get(\"f\", None)\n    if isinstance(path, str) and path.endswith((\".pt\", \".ckpt\", \".pth\", \".bin\", \".safetensors\")):\n        if path not in _ckpt_cache:\n            print(f\"  [cache] Loading checkpoint: {os.path.basename(path)}\")\n            _ckpt_cache[path] = _real_torch_load(*args, **kwargs)\n        else:\n            print(f\"  [cache] Reusing cached checkpoint: {os.path.basename(path)}\")\n        return _ckpt_cache[path]\n    return _real_torch_load(*args, **kwargs)\n\ntorch.load = _cached_torch_load\n\n# ── Import boltz + CliRunner ONCE ────────────────────────────────────────────\nfrom click.testing import CliRunner\nfrom boltz.main import cli as boltz_cli\n\n# ═════════════════════════════════════════════════════════════════════════════\n# Utilities\n# ═════════════════════════════════════════════════════════════════════════════\n\ndef extract_boltz_c1_coords(pdb_path, seq_len):\n    try:\n        from biopandas.pdb import PandasPdb\n        pl = PandasPdb().read_pdb(pdb_path)\n        res_coords = pl.df[\"ATOM\"][\n            (pl.df[\"ATOM\"].chain_id == \"0\")\n            & (pl.df[\"ATOM\"].atom_name == \"C1'\")\n        ]\n        if len(res_coords) == 0:\n            return None\n        coords = res_coords[[\"x_coord\", \"y_coord\", \"z_coord\"]].values.astype(np.float64)\n        if coords.shape[0] < seq_len:\n            coords = np.concatenate(\n                [coords, np.zeros((seq_len - coords.shape[0], 3), dtype=np.float64)]\n            )\n        return coords[:seq_len]\n    except Exception as e:\n        print(f\"    Error parsing {pdb_path}: {e}\")\n        return None\n\n# ═════════════════════════════════════════════════════════════════════════════\n# Core prediction via CliRunner (in-process, no subprocess)\n# ═════════════════════════════════════════════════════════════════════════════\n\ndef run_boltz_predict(fasta_path, out_dir, cache_dir, seed, cli_runner):\n    \"\"\"Invoke `boltz predict` in the SAME process via CliRunner.\"\"\"\n    cli_args = [\n        \"predict\", fasta_path,\n        \"--num_workers\", \"1\",\n        \"--max_parallel_samples\", \"1\",\n        \"--output_format\", \"pdb\",\n        \"--cache\", cache_dir,\n        \"--out_dir\", out_dir,\n        \"--seed\", str(seed),\n    ]\n    try:\n        result = cli_runner.invoke(boltz_cli, cli_args, catch_exceptions=True)\n        if result.exit_code != 0:\n            print(f\"    Boltz2 CLI exit code {result.exit_code}\")\n            if result.output:\n                for line in result.output.strip().splitlines()[-5:]:\n                    print(f\"      {line}\")\n            return False\n        return True\n    except Exception as e:\n        print(f\"    Boltz2 CliRunner exception: {e}\")\n        traceback.print_exc()\n        return False\n\ndef process_job(job, cli_runner, cache_dir, work_dir):\n    \"\"\"Process one (target, repeat) job.  Returns (tid, repeat, coords|None).\"\"\"\n    tid = job[\"target_id\"]\n    repeat = job[\"repeat\"]\n    seq_len = job[\"seq_len\"]\n    seed = job[\"seed\"]\n    fasta_path = job[\"fasta_path\"]\n\n    out_dir = os.path.join(work_dir, f\"out_{tid}_{repeat}\")\n    ok = run_boltz_predict(fasta_path, out_dir, cache_dir, seed, cli_runner)\n\n    coord = None\n    if ok:\n        pdb_path = (f\"{out_dir}/boltz_results_{tid}/predictions/\"\n                    f\"{tid}/{tid}_model_0.pdb\")\n        coord = extract_boltz_c1_coords(pdb_path, seq_len)\n\n    shutil.rmtree(out_dir, ignore_errors=True)\n    return (tid, repeat, coord)\n\n# ═════════════════════════════════════════════════════════════════════════════\n# Main entry point\n# ═════════════════════════════════════════════════════════════════════════════\n\ndef main():\n    parser = argparse.ArgumentParser()\n    parser.add_argument(\"--jobs_pkl\", required=True)\n    parser.add_argument(\"--results_pkl\", required=True)\n    parser.add_argument(\"--cache_dir\", required=True)\n    parser.add_argument(\"--work_dir\", required=True)\n    args = parser.parse_args()\n\n    os.makedirs(args.work_dir, exist_ok=True)\n\n    with open(args.jobs_pkl, \"rb\") as f:\n        jobs = pickle.load(f)\n\n    pid = os.getpid()\n    print(f\"[PID {pid}] Boltz2 worker started — {len(jobs)} job(s)\")\n    print(f\"[PID {pid}] CUDA_VISIBLE_DEVICES={os.environ.get('CUDA_VISIBLE_DEVICES')}\")\n\n    cli_runner = CliRunner()\n    results = []\n\n    for i, job in enumerate(jobs):\n        tid = job[\"target_id\"]\n        rep = job[\"repeat\"]\n        print(f\"\\n  [{i+1}/{len(jobs)}] {tid} r{rep}\")\n\n        try:\n            result = process_job(job, cli_runner, args.cache_dir, args.work_dir)\n            results.append(result)\n            status = \"OK\" if result[2] is not None else \"FAILED\"\n            print(f\"  [{i+1}/{len(jobs)}] {tid} r{rep} -> {status}\")\n        except Exception as e:\n            print(f\"  [{i+1}/{len(jobs)}] {tid} r{rep} -> EXCEPTION: {e}\")\n            traceback.print_exc()\n            results.append((tid, rep, None))\n\n        gc.collect()\n        torch.cuda.empty_cache()\n        gc.collect()\n\n    with open(args.results_pkl, \"wb\") as f:\n        pickle.dump(results, f)\n\n    print(f\"\\n[PID {pid}] Done — {len(results)} result(s) saved\")\n\nif __name__ == \"__main__\":\n    main()\n'''\n\n    path = \"/kaggle/working/boltz_gpu_worker.py\"\n    with open(path, \"w\", encoding=\"utf-8\") as f:\n        f.write(worker_code)\n    return path\n\ndef _extract_boltz_c1_coords(pdb_path, seq_len):\n    \"\"\"Kept for any non-worker usage.\"\"\"\n    try:\n        from biopandas.pdb import PandasPdb\n        pl = PandasPdb().read_pdb(pdb_path)\n        res_coords = pl.df[\"ATOM\"][\n            (pl.df[\"ATOM\"].chain_id == \"0\")\n            & (pl.df[\"ATOM\"].atom_name == \"C1'\")\n        ]\n        if len(res_coords) == 0:\n            return None\n        coords = res_coords[[\"x_coord\", \"y_coord\", \"z_coord\"]].values.astype(np.float64)\n        if coords.shape[0] < seq_len:\n            coords = np.concatenate(\n                [coords, np.zeros((seq_len - coords.shape[0], 3), dtype=np.float64)]\n            )\n        return coords[:seq_len]\n    except Exception as e:\n        print(f\"    Error parsing {pdb_path}: {e}\")\n        return None\n\ndef _write_boltz_fasta(target_id, row, fasta_dir, sequence_override=None):\n    \"\"\"Write multi-chain FASTA for a target (unchanged).\"\"\"\n    os.makedirs(fasta_dir, exist_ok=True)\n    fasta_path = os.path.join(fasta_dir, f\"{target_id}.fasta\")\n\n    if sequence_override is not None:\n        with open(fasta_path, \"w\") as fout:\n            fout.write(f\">0|rna|\\n{sequence_override}\\n\")\n        return fasta_path\n\n    pred_seq = row[\"sequence\"]\n    nucleotides = {\"A\", \"G\", \"C\", \"U\"}\n    aminoacids = {\"A\", \"R\", \"N\", \"D\", \"C\", \"E\", \"Q\", \"G\", \"H\", \"I\",\n                  \"L\", \"K\", \"M\", \"F\", \"P\", \"S\", \"T\", \"W\", \"Y\", \"V\"}\n    max_repeats = 999\n\n    try:\n        all_seq_raw = row.get(\"all_sequences\", \"\")\n        stoich_raw = row.get(\"stoichiometry\", \"\")\n        if pd.isna(all_seq_raw) or str(all_seq_raw).strip() == \"\":\n            raise ValueError(\"no all_sequences\")\n\n        chains, chain_ind = {}, 1\n        for entry in str(all_seq_raw).split(\"\\n\"):\n            entry = entry.strip()\n            if not entry:\n                continue\n            if entry.startswith(\">\"):\n                parts = entry.split(\"|\")\n                repeats = min(len(parts[1].split(\",\")) if len(parts) > 1 else 1,\n                              max_repeats)\n            else:\n                if entry == pred_seq:\n                    chains[0] = (entry, \"rna\")\n                else:\n                    entry_type = (\"rna\" if set(entry) <= nucleotides\n                                  else \"protein\" if set(entry) <= aminoacids\n                                  else None)\n                    if entry_type is None:\n                        continue\n                    chains[chain_ind] = (entry, entry_type)\n                    chain_ind += 1\n                    for _ in range(1, repeats):\n                        chains[chain_ind] = (entry, entry_type)\n                        chain_ind += 1\n\n        from random import shuffle as _shuffle\n        filtered, ntokens = {}, 0\n        other_keys = [c for c in chains if c != 0]\n        _shuffle(other_keys)\n        for c in [0] + other_keys:\n            ntokens += len(chains[c][0])\n            if ntokens >= BOLTZ_MAX_TOKENS:\n                break\n            filtered[c] = chains[c]\n        chains = filtered\n\n    except Exception:\n        chains = {0: (pred_seq, \"rna\")}\n\n    with open(fasta_path, \"w\") as fout:\n        for chain in sorted(chains.keys()):\n            seq, seq_type = chains[chain]\n            msa = \"\" if seq_type == \"rna\" else \"empty\"\n            fout.write(f\">{chain}|{seq_type}|{msa}\\n{seq}\\n\")\n\n    return fasta_path\n\ndef boltz2_phase(\n    test_df,\n    boltz_queue,\n    segments_map,\n    gpu_ids=None,\n):\n    n_gpus = _get_n_gpus()\n    if gpu_ids is None:\n        gpu_ids = list(range(n_gpus))\n\n    print(f\"\\n{'=' * 60}\")\n    print(f\"PHASE 2.5: Boltz2  (no chunking, persistent workers)\")\n    print(f\"  BOLTZ_MAX_TOKENS = {BOLTZ_MAX_TOKENS}\")\n    print(f\"  GPUs             = {gpu_ids}\")\n    print(f\"{'=' * 60}\")\n    t0 = time.time()\n\n    boltz_preds: dict = {}\n\n    # ── Collect all jobs (one per target × repeat, NO chunking) ───────────\n    all_jobs = []\n    target_meta = {}\n\n    os.makedirs(BOLTZ_INPUT_DIR, exist_ok=True)\n\n    for _, row in test_df.iterrows():\n        tid = row[\"target_id\"]\n        seq = row[\"sequence\"]\n        seq_len = len(seq)\n        n_boltz_needed = boltz_queue.get(tid, (0, \"\"))[0]\n\n        if n_boltz_needed <= 0:\n            print(f\"  {tid} ({seq_len} nt): 0 Boltz2 slots — skip\")\n            continue\n\n        fasta_path = _write_boltz_fasta(tid, row, BOLTZ_INPUT_DIR)\n        target_meta[tid] = {\"seq_len\": seq_len, \"n_needed\": n_boltz_needed}\n\n        print(f\"  {tid} ({seq_len} nt): need {n_boltz_needed} from Boltz2\")\n\n        for repeat in range(n_boltz_needed):\n            repeat_seed = SEED + repeat * 1000 + stable_hash(tid) % 10000\n            all_jobs.append({\n                \"target_id\": tid,\n                \"repeat\": repeat,\n                \"full_sequence\": seq,\n                \"seq_len\": seq_len,\n                \"seed\": repeat_seed,\n                \"fasta_path\": fasta_path,\n            })\n\n    if not all_jobs:\n        print(\"\\n  No targets need Boltz2\")\n        elapsed = time.time() - t0\n        print(f\"\\nPhase 2.5 done in {elapsed:.1f}s\")\n        return boltz_preds\n\n    # ── Partition jobs across available GPUs (round-robin) ────────────────\n    effective_gpus = min(len(gpu_ids), len(all_jobs))\n    active_gpu_ids = gpu_ids[:effective_gpus]\n    gpu_jobs = [[] for _ in range(effective_gpus)]\n    for i, job in enumerate(all_jobs):\n        gpu_jobs[i % effective_gpus].append(job)\n\n    for slot in range(effective_gpus):\n        tids = [j[\"target_id\"] for j in gpu_jobs[slot]]\n        print(f\"  GPU {active_gpu_ids[slot]}: {len(gpu_jobs[slot])} job(s) → {tids}\")\n\n    # ── Write worker script + per-GPU job pickles ─────────────────────────\n    worker_py = _write_boltz_gpu_worker_code()\n\n    job_pkl_paths = []\n    result_pkl_paths = []\n    work_dirs = []\n\n    for slot in range(effective_gpus):\n        gid = active_gpu_ids[slot]\n        job_pkl = f\"/kaggle/working/boltz_jobs_gpu{gid}.pkl\"\n        result_pkl = f\"/kaggle/working/boltz_results_gpu{gid}.pkl\"\n        work_dir = f\"/kaggle/working/boltz_work_gpu{gid}\"\n\n        with open(job_pkl, \"wb\") as f:\n            _pickle.dump(gpu_jobs[slot], f)\n\n        job_pkl_paths.append(job_pkl)\n        result_pkl_paths.append(result_pkl)\n        work_dirs.append(work_dir)\n\n    # ── Launch ONE subprocess per GPU ─────────────────────────────────────\n    print(f\"\\n  Launching {effective_gpus} persistent Boltz2 worker(s)...\")\n\n    procs = []\n    for slot in range(effective_gpus):\n        if not gpu_jobs[slot]:\n            continue\n        gid = active_gpu_ids[slot]\n\n        env = os.environ.copy()\n        env[\"CUDA_VISIBLE_DEVICES\"] = str(gid)\n        env[\"PYTHONHASHSEED\"] = str(SEED)\n        env[\"CUBLAS_WORKSPACE_CONFIG\"] = \":4096:8\"\n        env[\"NVIDIA_TF32_OVERRIDE\"] = \"0\"\n\n        cmd = [\n            sys.executable, worker_py,\n            \"--jobs_pkl\",      job_pkl_paths[slot],\n            \"--results_pkl\",   result_pkl_paths[slot],\n            \"--cache_dir\",     BOLTZ_CACHE_DIR,\n            \"--work_dir\",      work_dirs[slot],\n        ]\n\n        p = _subprocess.Popen(\n            cmd, stdout=_subprocess.PIPE, stderr=_subprocess.STDOUT,\n            env=env, text=True,\n        )\n        procs.append((gid, p))\n        print(f\"  GPU {gid}: PID {p.pid} started \"\n              f\"({len(gpu_jobs[slot])} jobs)\")\n\n    # ── Wait and stream output ────────────────────────────────────────────\n    for gid, p in procs:\n        stdout, _ = p.communicate()\n        if stdout:\n            for line in stdout.splitlines():\n                print(f\"  [GPU{gid}] {line}\")\n        if p.returncode != 0:\n            print(f\"  ⚠ GPU {gid} Boltz2 worker returned \"\n                  f\"exit code {p.returncode}\")\n\n    # ── Parse results from all GPU workers ────────────────────────────────\n    from collections import defaultdict\n    coords_by_target = defaultdict(list)\n\n    for slot in range(effective_gpus):\n        result_pkl = result_pkl_paths[slot]\n        if not os.path.isfile(result_pkl):\n            print(f\"  ⚠ Missing result pickle: {result_pkl}\")\n            continue\n        with open(result_pkl, \"rb\") as f:\n            results = _pickle.load(f)\n        for (tid, repeat, coord) in results:\n            if coord is not None:\n                coords_by_target[tid].append(coord)\n            else:\n                print(f\"  {tid} repeat {repeat}: no valid C1' coords\")\n\n    for tid in target_meta:\n        coords = coords_by_target.get(tid, [])\n        if coords:\n            boltz_preds[tid] = coords\n            print(f\"  [{tid}] Boltz2 produced {len(coords)} structure(s)\")\n        else:\n            print(f\"  [{tid}] Boltz2 produced 0 valid structures\")\n\n    # ── Cleanup ───────────────────────────────────────────────────────────\n    if os.path.exists(BOLTZ_INPUT_DIR):\n        shutil.rmtree(BOLTZ_INPUT_DIR, ignore_errors=True)\n    for p in job_pkl_paths + result_pkl_paths:\n        try: os.remove(p)\n        except OSError: pass\n    for d in work_dirs:\n        shutil.rmtree(d, ignore_errors=True)\n    try: os.remove(worker_py)\n    except OSError: pass\n\n    gc.collect()\n    torch.cuda.empty_cache()\n    gc.collect()\n\n    elapsed = time.time() - t0\n    print(f\"\\nPhase 2.5 done in {elapsed:.1f}s  (GPUs {active_gpu_ids})\")\n    return boltz_preds","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_combined_preds(test_df, template_preds, protenix_preds, boltz_preds, segments_map):\n    combined = {}\n    for _, row in test_df.iterrows():\n        tid = row[\"target_id\"]\n        seq = row[\"sequence\"]\n\n        preds = list(template_preds.get(tid, []))\n\n        # Append Protenix samples\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(preds) >= N_SAMPLE:\n                    break\n                preds.append(adaptive_rna_constraints(\n                    np.asarray(ptx[j], np.float64), tid, segments_map,\n                    confidence=0.55, passes=1))\n\n        # Append Boltz2 samples as direct slots\n        for bc in boltz_preds.get(tid, []):\n            if len(preds) >= N_SAMPLE:\n                break\n            preds.append(adaptive_rna_constraints(\n                np.asarray(bc, np.float64), tid, segments_map,\n                confidence=0.55, passes=1))\n\n        # De-novo fallback\n        while len(preds) < N_SAMPLE:\n            seed_val = stable_hash(tid) % 10000 + len(preds) * 1000\n            dn = generate_rna_structure(seq, seed=seed_val)\n            preds.append(adaptive_rna_constraints(dn, tid, segments_map, confidence=0.2))\n\n        combined[tid] = preds[:N_SAMPLE]\n    return combined","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ═══════════════════════════════════════════════════════════════════════════════\n# §14  Save Submission\n# ═══════════════════════════════════════════════════════════════════════════════\n\ndef save_submission(test_df, combined_preds, output_csv):\n    \"\"\"Phase 4 — Write final combined_preds to submission CSV.\"\"\"\n    print(f\"\\n{'=' * 60}\")\n    print(\"PHASE 4: Save Submission\")\n    print(f\"{'=' * 60}\")\n\n    all_rows = []\n    for _, row in test_df.iterrows():\n        tid = row[\"target_id\"]\n        seq = row[\"sequence\"]\n        preds = combined_preds.get(tid, [])\n\n        # Safety: pad if somehow short\n        while len(preds) < N_SAMPLE:\n            seed_val = stable_hash(tid) % 10000 + len(preds) * 1000\n            preds.append(generate_rna_structure(seq, seed=seed_val))\n\n        print(f\"  {tid}: {len(preds)} predictions\")\n\n        stacked = np.stack(preds[: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\"] + [\n        f\"{c}_{i}\" for i in range(1, N_SAMPLE + 1) for c in [\"x\", \"y\", \"z\"]\n    ]\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)\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ═══════════════════════════════════════════════════════════════════════════════\n# main()  — UPDATED training data loading section\n#   Replace the training data loading + tbm_phase call in main()\n# ═══════════════════════════════════════════════════════════════════════════════\n\ndef main() -> None:\n    test_csv, output_csv, code_dir, root_dir = resolve_paths()\n\n    if not os.path.isdir(code_dir):\n        raise FileNotFoundError(\n            f\"Missing PROTENIX_CODE_DIR: {code_dir}. \"\n            \"Set PROTENIX_CODE_DIR to the repo path.\"\n        )\n\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    test_df_full = pd.read_csv(test_csv)\n    test_df = (\n        test_df_full.head(LOCAL_N_SAMPLES) if not IS_KAGGLE else test_df_full\n    ).reset_index(drop=True)\n    print(f\"Test targets : {len(test_df)}\"\n          + (\" (LOCAL MODE)\" if not IS_KAGGLE else \"\"))\n\n    test_df_trunc = test_df.copy()\n    test_df_trunc[\"sequence\"] = test_df_trunc[\"sequence\"].str[:MAX_SEQ_LEN]\n\n    # ── Build training database (NEW: chain-aware) ────────────────────────\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, low_memory=False)\n    val_labels = pd.read_csv(DEFAULT_VAL_LBLS, low_memory=False)\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\n    full_coords, chain_index = build_training_database(combined_seqs, combined_labels)\n    segments_map, _ = build_segments_map(test_df)\n\n    print(f\"Template pool: {len(chain_index)} entries, \"\n          f\"{len(full_coords)} structures\")\n\n    # ── PHASE 1: TBM (NEW: pooled FULL + CHAIN) ──────────────────────────\n    template_preds, protenix_queue, boltz_queue = tbm_phase(\n        test_df, full_coords, chain_index, segments_map\n    )\n\n    # ── PHASE 2: Protenix (both GPUs) ────────────────────────────────────\n    protenix_preds: dict = {}\n\n    if protenix_queue and USE_PROTENIX:\n        protenix_preds = protenix_phase(\n            protenix_queue, test_df_trunc,\n            lambda p: build_configs(\n                p, str(Path(\"/kaggle/working\") / \"outputs\"), MODEL_NAME\n            ),\n            template_preds=template_preds,\n        )\n    elif protenix_queue and not USE_PROTENIX:\n        print(f\"\\nPHASE 2 skipped (USE_PROTENIX=False).\")\n\n    # ── PHASE 2.5: Boltz2 (both GPUs, after Protenix is done) ────────────\n    boltz_preds: dict = {}\n\n    want_boltz = USE_BOLTZ2 and any(\n        boltz_queue.get(tid, (0, \"\"))[0] > 0\n        for tid in boltz_queue\n    )\n\n    if want_boltz:\n        try:\n            boltz_preds = boltz2_phase(\n                test_df, boltz_queue, segments_map,\n            )\n        except Exception as e:\n            print(f\"  ⚠ Boltz2 phase failed: {e}\")\n            import traceback; traceback.print_exc()\n\n    # ── Build combined predictions (TBM + Protenix + Boltz2 as slots) ─────\n    combined_preds = build_combined_preds(\n        test_df, template_preds, protenix_preds, boltz_preds, segments_map\n    )\n    for tid, preds in combined_preds.items():\n        seq_len = len(test_df[test_df[\"target_id\"] == tid][\"sequence\"].values[0])\n        n_tbm = len(template_preds.get(tid, []))\n        ptx = protenix_preds.get(tid)\n        n_ptx = ptx.shape[0] if ptx is not None and ptx.ndim == 3 else 0\n        n_boltz = min(len(boltz_preds.get(tid, [])), N_SAMPLE - n_tbm - n_ptx)\n        n_dn = N_SAMPLE - n_tbm - min(n_ptx, N_SAMPLE - n_tbm) - max(n_boltz, 0)\n        print(f\"  {tid} ({seq_len} nt): combined = \"\n              f\"tbm({n_tbm}) + ptx({min(n_ptx, N_SAMPLE - n_tbm)}) \"\n              f\"+ boltz({max(n_boltz, 0)}) + denovo({max(n_dn, 0)})\")\n\n    # ── PHASE 3: RNAPro refinement on ALL 5 combined predictions ──────────\n    if USE_RNAPRO:\n        try:\n            rnapro_refine_all(test_df, combined_preds)\n        except Exception as e:\n            print(f\"  ⚠ RNAPro refinement failed: {e}\")\n            import traceback\n            traceback.print_exc()\n            print(\"  Continuing with unrefined combined predictions\")\n\n    # ── PHASE 4: Save ─────────────────────────────────────────────────────\n    save_submission(test_df, combined_preds, output_csv)\n\n    # ── Cleanup ──────────────────────────────────────────────────────────\n    for item in os.listdir(\"/kaggle/working\"):\n        path = os.path.join(\"/kaggle/working\", item)\n        if item == \"submission.csv\":\n            print(f\"  KEEP: {item}\"); continue\n        try:\n            if os.path.isdir(path): shutil.rmtree(path)\n            else: os.remove(path)\n            print(f\"  DEL:  {item}\")\n        except Exception as e:\n            print(f\"  SKIP: {item} ({e})\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}