{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":118765,"databundleVersionId":15231210,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":10880374,"sourceType":"datasetVersion","datasetId":6760482},{"sourceId":10880419,"sourceType":"datasetVersion","datasetId":6760509},{"sourceId":11230242,"sourceType":"datasetVersion","datasetId":7014687},{"sourceId":11451236,"sourceType":"datasetVersion","datasetId":7174725},{"sourceId":11899194,"sourceType":"datasetVersion","datasetId":7479946},{"sourceId":13282339,"sourceType":"datasetVersion","datasetId":7162026}],"dockerImageVersionId":31259,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\"\"\"\nTemplate + Protenix Combined RNA 3D Folding (script version)\n\nWhat this script does:\n- Detects Kaggle vs local environment and sets paths\n- (Kaggle only) installs Protenix/USalign wheel deps and prepares Protenix repo\n- Loads train/val/test sequences + labels\n- Phase 1: template-based predictions (fast alignment + coord adaptation + constraints)\n- Phase 2: Protenix fallback for targets missing template slots\n- Phase 3: merges into 5 predictions per target and writes submission.csv\n- (Optional) runs validation + TM-score evaluation via USalign\n\"\"\"\n\n# ---------------------------\n# Imports (standard library)\n# ---------------------------\nimport contextlib\nimport json\nimport os\nimport random\nimport re\nimport shutil\nimport subprocess\nimport sys\nimport time\nimport warnings\nfrom pathlib import Path\n\n# ---------------------------\n# Imports (third-party)\n# ---------------------------\nimport numpy as np\nimport pandas as pd\n\nwarnings.filterwarnings(\"ignore\")\n\n# ===========================================================================\n# Hardcoded config (Kaggle-friendly, no argparse)\n# ===========================================================================\nCONFIG = {\n    # Paths\n    # Set to None to auto-detect (/kaggle/input/... or /data/kaggle/...)\n    \"DATA_PATH\": None,\n\n    # Main toggles\n    \"SHOW_VALIDATION\": 0,\n    \"MAKE_SUBMISSION\": 1,\n    \"USE_PROTENIX\": 1,\n\n    # Template thresholds\n    \"MIN_SIMILARITY\": 0.0,\n    \"MIN_PERCENT_IDENTITY\": 50.0,\n\n    # Debug / verbosity\n    \"DEBUG\": 0,\n    \"SHOW_ALIGNMENT_DETAILS\": 1,\n}\n\n# ===========================================================================\n# Kaggle dataset paths (fixed)\n# ===========================================================================\nKAGGLE_DATASETS = {\n    \"PROTENIX_PACKAGES\": \"/kaggle/input/datasets/zoushuxian/protenix-packages/packages\",\n    \"PROTENIX_CHECKPOINT\": \"/kaggle/input/datasets/zoushuxian/protenix-finetuned-rna3db-all-1599/1599_ema_0.999.pt\",\n    \"PROTENIX_REPO\": \"/kaggle/input/datasets/zoushuxian/protenix-rmsa-repo/protenix_kaggle\",\n    \"BIOPYTHON_WHL\": \"/kaggle/input/datasets/ogurtsov/biopython/biopython-1.85-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\",\n    \"ML_COLLECTIONS_WHL\": \"/kaggle/input/datasets/ogurtsov/ml-collections/ml_collections-1.0.0-py3-none-any.whl\",\n}\n\n\n# ---------------------------\n# Utility: seeding\n# ---------------------------\ndef seed_everything(seed: int = 42):\n    random.seed(seed)\n    np.random.seed(seed)\n\n\n# ---------------------------\n# Utility: shell commands\n# ---------------------------\ndef run_cmd(cmd: str, cwd: str | None = None, check: bool = True):\n    proc = subprocess.run(\n        cmd,\n        shell=True,\n        cwd=cwd,\n        text=True,\n        capture_output=True,\n    )\n\n    if check and proc.returncode != 0:\n        raise RuntimeError(\n            f\"Command failed (code={proc.returncode}): {cmd}\\n\\nSTDOUT:\\n{proc.stdout}\\n\\nSTDERR:\\n{proc.stderr}\"\n        )\n\n    return proc.stdout + proc.stderr\n\n\n# ---------------------------\n# Environment detection + config\n# ---------------------------\ndef detect_environment(user_data_path: str | None):\n    # Prefer explicit user path\n    if user_data_path is not None and os.path.exists(user_data_path):\n        is_kaggle = user_data_path.startswith(\"/kaggle/\")\n        data_path = user_data_path\n        return is_kaggle, data_path\n\n    # Kaggle default\n    if os.path.exists(\"/kaggle/input/stanford-rna-3d-folding-2/\"):\n        return True, \"/kaggle/input/stanford-rna-3d-folding-2/\"\n\n    # Common local mount\n    if os.path.exists(\"/data/kaggle/stanford-rna-3d-folding-2/\"):\n        return False, \"/data/kaggle/stanford-rna-3d-folding-2/\"\n\n    raise ValueError(\"Data path not found. Provide DATA_PATH explicitly in CONFIG.\")\n\n\ndef build_paths(is_kaggle: bool, data_path: str):\n    if is_kaggle:\n        output_path = \"/kaggle/working/output\"\n        usalign_bin = \"/kaggle/working/USalign\"\n        protenix_dir = \"/kaggle/working/Protenix\"\n    else:\n        output_path = \"./output\"\n        usalign_bin = \"./USalign/USalign\"\n        protenix_dir = f\"{data_path}/protenix_kaggle\"\n\n    return {\n        \"IS_KAGGLE\": is_kaggle,\n        \"DATA_PATH\": data_path if data_path.endswith(\"/\") else data_path + \"/\",\n        \"OUTPUT_PATH\": output_path,\n        \"USALIGN_BIN\": usalign_bin,\n        \"PROTENIX_DIR\": protenix_dir,\n    }\n\n\n# ---------------------------\n# Kaggle-only installs / setup\n# ---------------------------\ndef kaggle_setup(paths: dict):\n    # Copy USalign binary\n    run_cmd(f\"cp {KAGGLE_DATASETS['PROTENIX_PACKAGES']}/USalign /kaggle/working/\")\n    run_cmd(\"chmod +x /kaggle/working/USalign\")\n\n    # Add utils path\n    sys.path.insert(0, \"/kaggle/input/rna-3d-utils/\")\n\n    # Install Protenix wheels (no-deps)\n    run_cmd(f\"cp -r {KAGGLE_DATASETS['PROTENIX_PACKAGES']} /kaggle/working/packages\")\n    run_cmd(\"pip install --no-deps --exists-action=i *.whl\", cwd=\"/kaggle/working/packages\")\n\n    # Install ihm/modelcif from unpacked folders (as in notebook)\n    run_cmd(\"mv /kaggle/working/packages/ihm-2.3/ihm-2.3 /kaggle/working\")\n    run_cmd(\"mv /kaggle/working/packages/modelcif-0.7/modelcif-0.7 /kaggle/working\")\n\n    run_cmd(\"pip install /kaggle/working/ihm-2.3\")\n    run_cmd(\"pip install /kaggle/working/modelcif-0.7\")\n\n    run_cmd(\"rm -rf /kaggle/working/ihm-2.3\")\n    run_cmd(\"rm -rf /kaggle/working/modelcif-0.7\")\n\n    # Biopython + ml-collections wheels\n    run_cmd(f\"pip install {KAGGLE_DATASETS['BIOPYTHON_WHL']}\")\n    run_cmd(f\"pip install {KAGGLE_DATASETS['ML_COLLECTIONS_WHL']}\")\n\n    run_cmd(\"rm -rf /kaggle/working/packages\")\n\n    # Install protenix_mg_packages wheels (no-deps)\n    run_cmd(\"cp -r /kaggle/input/protenix-mg-packages/protenix_mg_packages /kaggle/working\")\n    run_cmd(\"pip install --no-deps --exists-action=i *.whl\", cwd=\"/kaggle/working/protenix_mg_packages\")\n    run_cmd(\"rm -rf /kaggle/working/protenix_mg_packages\")\n\n    # Copy Protenix repo folder into /kaggle/working/Protenix\n    run_cmd(f\"cp -R {KAGGLE_DATASETS['PROTENIX_REPO']} /kaggle/working/\")\n    run_cmd(\"mv /kaggle/working/protenix_kaggle /kaggle/working/Protenix\")\n\n\n# ---------------------------\n# Load dataset CSVs\n# ---------------------------\ndef load_sequence_and_label_data(data_path: str):\n    print(\"Loading sequence data...\")\n\n    train_seqs = pd.read_csv(data_path + \"train_sequences.csv\")\n    validation_seqs = pd.read_csv(data_path + \"validation_sequences.csv\")\n    test_seqs = pd.read_csv(data_path + \"test_sequences.csv\")\n\n    train_labels = pd.read_csv(data_path + \"train_labels.csv\")\n    validation_labels = pd.read_csv(data_path + \"validation_labels.csv\")\n\n    print(f\"Loaded {len(train_seqs)} training sequences\")\n    print(f\"Loaded {len(validation_seqs)} validation sequences\")\n    print(f\"Loaded {len(test_seqs)} test sequences\")\n\n    return train_seqs, validation_seqs, test_seqs, train_labels, validation_labels\n\n\n# ===========================================================================\n# Template-based modelling\n# ===========================================================================\ndef make_aligner():\n    # Import here so script can still parse even if BioPython missing locally\n    from Bio.Align import PairwiseAligner\n\n    al = PairwiseAligner()\n    al.mode = \"global\"\n    al.match_score = 2\n    al.mismatch_score = -1.5\n\n    al.open_gap_score = -8\n    al.extend_gap_score = -0.4\n\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\n    return al\n\n\ndef parse_fasta(fasta_content: str):\n    out = {}\n    cur = None\n    seq_parts = []\n\n    for line in str(fasta_content).splitlines():\n        line = line.strip()\n        if not line:\n            continue\n\n        if line.startswith(\">\"):\n            if cur is not None:\n                out[cur] = \"\".join(seq_parts)\n            cur = line[1:].split()[0]\n            seq_parts = []\n        else:\n            seq_parts.append(line.replace(\" \", \"\"))\n\n    if cur is not None:\n        out[cur] = \"\".join(seq_parts)\n\n    return out\n\n\ndef parse_stoichiometry(stoich: str):\n    if pd.isna(stoich) or str(stoich).strip() == \"\":\n        return []\n\n    out = []\n    for part in str(stoich).split(\";\"):\n        ch, cnt = part.split(\":\")\n        out.append((ch.strip(), int(cnt)))\n\n    return out\n\n\ndef get_chain_segments(row: pd.Series):\n    seq = row[\"sequence\"]\n    stoich = row.get(\"stoichiometry\", \"\")\n    all_seq = row.get(\"all_sequences\", \"\")\n\n    if (\n        pd.isna(stoich)\n        or pd.isna(all_seq)\n        or str(stoich).strip() == \"\"\n        or str(all_seq).strip() == \"\"\n    ):\n        return [(0, len(seq))]\n\n    try:\n        chain_dict = parse_fasta(all_seq)\n        order = parse_stoichiometry(stoich)\n        segs = []\n        pos = 0\n\n        for ch, cnt in order:\n            base = chain_dict.get(ch)\n            if base is None:\n                return [(0, len(seq))]\n\n            for _ in range(cnt):\n                L = len(base)\n                segs.append((pos, pos + L))\n                pos += L\n\n        if pos != len(seq):\n            return [(0, len(seq))]\n\n        return segs\n    except Exception:\n        return [(0, len(seq))]\n\n\ndef build_segments_map(df: pd.DataFrame):\n    seg_map = {}\n    stoich_map = {}\n\n    for _, r in df.iterrows():\n        tid = r[\"target_id\"]\n        seg_map[tid] = get_chain_segments(r)\n        stoich_map[tid] = str(\n            r.get(\"stoichiometry\", \"\")\n            if not pd.isna(r.get(\"stoichiometry\", \"\"))\n            else \"\"\n        )\n\n    return seg_map, stoich_map\n\n\ndef process_labels(labels_df: pd.DataFrame):\n    coords_dict = {}\n    prefixes = labels_df[\"ID\"].str.rsplit(\"_\", n=1).str[0]\n\n    for id_prefix, group in labels_df.groupby(prefixes):\n        coords_dict[id_prefix] = group.sort_values(\"resid\")[[\"x_1\", \"y_1\", \"z_1\"]].values\n\n    return coords_dict\n\n\ndef _build_aligned_strings(query_seq: str, template_seq: str, alignment):\n    q_segments, t_segments = alignment.aligned\n    aligned_q = []\n    aligned_t = []\n    qi = 0\n    ti = 0\n\n    for (qs, qe), (ts, te) in zip(q_segments, t_segments):\n        while qi < qs:\n            aligned_q.append(query_seq[qi])\n            aligned_t.append(\"-\")\n            qi += 1\n\n        while ti < ts:\n            aligned_q.append(\"-\")\n            aligned_t.append(template_seq[ti])\n            ti += 1\n\n        for q_pos, t_pos in zip(range(qs, qe), range(ts, te)):\n            aligned_q.append(query_seq[q_pos])\n            aligned_t.append(template_seq[t_pos])\n\n        qi = qe\n        ti = te\n\n    while qi < len(query_seq):\n        aligned_q.append(query_seq[qi])\n        aligned_t.append(\"-\")\n        qi += 1\n\n    while ti < len(template_seq):\n        aligned_q.append(\"-\")\n        aligned_t.append(template_seq[ti])\n        ti += 1\n\n    return \"\".join(aligned_q), \"\".join(aligned_t)\n\n\ndef find_similar_sequences_detailed(\n    query_seq: str,\n    train_seqs_df: pd.DataFrame,\n    train_coords_dict: dict,\n    aligner,\n    temporal_cutoff=None,\n    top_n: int = 5,\n):\n    similar_seqs = []\n\n    if temporal_cutoff is not None and \"temporal_cutoff\" in train_seqs_df.columns:\n        filtered = train_seqs_df[train_seqs_df[\"temporal_cutoff\"] < temporal_cutoff]\n    else:\n        filtered = train_seqs_df\n\n    for _, row in filtered.iterrows():\n        target_id = row[\"target_id\"]\n        train_seq = row[\"sequence\"]\n\n        if target_id not in train_coords_dict:\n            continue\n\n        if abs(len(train_seq) - len(query_seq)) / max(len(train_seq), len(query_seq)) > 0.3:\n            continue\n\n        alignment = next(iter(aligner.align(query_seq, train_seq)))\n        raw_score = alignment.score\n        normalized_score = raw_score / (2 * min(len(query_seq), len(train_seq)))\n\n        identical = 0\n        for (qs, qe), (ts, te) in zip(*alignment.aligned):\n            for q_pos, t_pos in zip(range(qs, qe), range(ts, te)):\n                if query_seq[q_pos] == train_seq[t_pos]:\n                    identical += 1\n\n        percent_identity = 100 * identical / len(query_seq)\n        aligned_query, aligned_template = _build_aligned_strings(query_seq, train_seq, alignment)\n\n        similar_seqs.append(\n            (\n                target_id,\n                train_seq,\n                normalized_score,\n                train_coords_dict[target_id],\n                percent_identity,\n                aligned_query,\n                aligned_template,\n            )\n        )\n\n    similar_seqs.sort(key=lambda x: x[2], reverse=True)\n    return similar_seqs[:top_n]\n\n\ndef adapt_template_to_query(query_seq: str, template_seq: str, template_coords: np.ndarray, aligner):\n    alignment = next(iter(aligner.align(query_seq, template_seq)))\n    new_coords = np.full((len(query_seq), 3), np.nan)\n\n    for (q_start, q_end), (t_start, t_end) in zip(*alignment.aligned):\n        t_chunk = template_coords[t_start:t_end]\n        if len(t_chunk) == (q_end - q_start):\n            new_coords[q_start:q_end] = t_chunk\n\n    for i in range(len(new_coords)):\n        if np.isnan(new_coords[i, 0]):\n            prev_v = next((j for j in range(i - 1, -1, -1) if not np.isnan(new_coords[j, 0])), -1)\n            next_v = next((j for j in range(i + 1, len(new_coords)) if not np.isnan(new_coords[j, 0])), -1)\n\n            if prev_v >= 0 and next_v >= 0:\n                w = (i - prev_v) / (next_v - prev_v)\n                new_coords[i] = (1 - w) * new_coords[prev_v] + w * new_coords[next_v]\n            elif prev_v >= 0:\n                new_coords[i] = new_coords[prev_v] + [3, 0, 0]\n            elif next_v >= 0:\n                new_coords[i] = new_coords[next_v] + [3, 0, 0]\n            else:\n                new_coords[i] = [i * 3, 0, 0]\n\n    return np.nan_to_num(new_coords)\n\n\ndef adaptive_rna_constraints(\n    coordinates: np.ndarray,\n    target_id: str,\n    segments_map: dict,\n    confidence: float = 1.0,\n    passes: int = 2,\n):\n    coords = coordinates.copy()\n    segments = segments_map.get(target_id, [(0, len(coords))])\n\n    strength = 0.75 * (1.0 - min(confidence, 0.97))\n    strength = max(strength, 0.02)\n\n    for _ in range(passes):\n        for (s, e) in segments:\n            X = coords[s:e]\n            L = e - s\n\n            if L < 3:\n                coords[s:e] = X\n                continue\n\n            # (1) bond i,i+1 to ~5.95A\n            d = X[1:] - X[:-1]\n            dist = np.linalg.norm(d, axis=1) + 1e-6\n            target = 5.95\n            scale = (target - dist) / dist\n            adj = (d * scale[:, None]) * (0.22 * strength)\n            X[:-1] -= adj\n            X[1:] += adj\n\n            # (2) soft i,i+2 to ~10.2A\n            d2 = X[2:] - X[:-2]\n            dist2 = np.linalg.norm(d2, axis=1) + 1e-6\n            target2 = 10.2\n            scale2 = (target2 - dist2) / dist2\n            adj2 = (d2 * scale2[:, None]) * (0.10 * strength)\n            X[:-2] -= adj2\n            X[2:] += adj2\n\n            # (3) Laplacian smoothing\n            lap = 0.5 * (X[:-2] + X[2:]) - X[1:-1]\n            X[1:-1] += (0.06 * strength) * lap\n\n            # (4) self-avoidance\n            if L >= 25:\n                k = min(L, 160) if L > 220 else L\n                if k < L:\n                    idx = np.linspace(0, L - 1, k).astype(int)\n                else:\n                    idx = np.arange(L)\n\n                P = X[idx]\n                diff = P[:, None, :] - P[None, :, :]\n                distm = np.linalg.norm(diff, axis=2) + 1e-6\n                sep = np.abs(idx[:, None] - idx[None, :])\n\n                mask = (sep > 2) & (distm < 3.2)\n                if np.any(mask):\n                    force = (3.2 - distm) / distm\n                    vec = (diff * force[:, :, None] * mask[:, :, None]).sum(axis=1)\n                    X[idx] += (0.015 * strength) * vec\n\n            coords[s:e] = X\n\n    return coords\n\n\ndef generate_rna_structure(sequence: str, seed: int | None = None):\n    if seed is not None:\n        np.random.seed(seed)\n\n    n = len(sequence)\n    coords = np.zeros((n, 3))\n\n    for i in range(n):\n        angle = i * 0.6\n        coords[i] = [10.0 * np.cos(angle), 10.0 * np.sin(angle), i * 2.5]\n\n    return coords\n\n\n# ===========================================================================\n# Protenix utilities\n# ===========================================================================\ndef extract_c1_atoms(cif_path: Path):\n    from biotite.structure.io.pdbx import CIFFile, get_structure\n\n    cif_file = CIFFile.read(str(cif_path))\n    model = get_structure(cif_file, model=1)\n    chain = model[model.chain_id == \"A\"]\n    mask = chain.atom_name == \"C1'\"\n    c1_atoms = chain[mask]\n\n    df = pd.DataFrame.from_dict(c1_atoms._annot)\n    df[\"x\"] = c1_atoms.coord[:, 0]\n    df[\"y\"] = c1_atoms.coord[:, 1]\n    df[\"z\"] = c1_atoms.coord[:, 2]\n\n    return df[[\"res_name\", \"res_id\", \"x\", \"y\", \"z\"]]\n\n\ndef prepare_protenix_json(target_id: str, sequence: str, output_path: Path, input_path: Path, max_length: int = 400):\n    if len(sequence) <= max_length:\n        input_json = [\n            {\n                \"sequences\": [\n                    {\n                        \"rnaSequence\": {\n                            \"sequence\": sequence,\n                            \"count\": 1,\n                            \"msa\": {\n                                \"precomputed_msa_dir\": f\"{input_path}/MSA/{target_id}.MSA.fasta\",\n                                \"pairing_db\": \"rnacentral\",\n                            },\n                        }\n                    }\n                ],\n                \"name\": target_id,\n            }\n        ]\n    else:\n        print(f\"    Sequence too long ({len(sequence)} > {max_length}), truncating, no MSA\")\n        input_json = [\n            {\n                \"sequences\": [\n                    {\n                        \"rnaSequence\": {\n                            \"sequence\": sequence[:max_length],\n                            \"count\": 1,\n                        }\n                    }\n                ],\n                \"name\": target_id,\n            }\n        ]\n\n    json_path = output_path / \"input_json\" / f\"{target_id}.json\"\n    json_path.parent.mkdir(parents=True, exist_ok=True)\n\n    with open(json_path, \"w\") as f:\n        json.dump(input_json, f, indent=4)\n\n\ndef run_protenix_inference(\n    is_kaggle: bool,\n    data_path: str,\n    target_id: str,\n    sequence: str,\n    output_path: Path,\n    input_path: Path,\n    protenix_dir: str,\n    seed: int = 101,\n    n_cycle: int = 10,\n    n_sample: int = 5,\n    n_step: int = 200,\n    max_length: int = 400,\n):\n    if is_kaggle:\n        checkpoint_path = KAGGLE_DATASETS[\"PROTENIX_CHECKPOINT\"]\n    else:\n        checkpoint_path = f\"{data_path}/protenix_chpt/1599_ema_0.999.pt\"\n\n    input_json_path = output_path / \"input_json\" / f\"{target_id}.json\"\n    dump_dir = output_path / target_id\n    dump_dir.mkdir(parents=True, exist_ok=True)\n\n    use_msa = \"True\" if len(sequence) <= max_length else \"False\"\n\n    # Runner expects argv\n    sys.argv = [\n        \"runner/inference.py\",\n        f\"--seeds={seed}\",\n        f\"--dump_dir={dump_dir}\",\n        f\"--input_json_path={input_json_path}\",\n        f\"--model.N_cycle={n_cycle}\",\n        f\"--sample_diffusion.N_sample={n_sample}\",\n        f\"--sample_diffusion.N_step={n_step}\",\n        \"--augment.use_rnalm True\",\n        f\"--use_msa {use_msa}\",\n        f\"--load_checkpoint_path={checkpoint_path}\",\n        \"\",\n    ]\n\n    with protenix_context(protenix_dir):\n        from runner.inference import run\n\n        run()\n\n\ndef get_protenix_predictions(target_id: str, sequence: str, output_path: Path, seed: int = 101, n_sample: int = 5):\n    predictions = []\n\n    for i in range(n_sample):\n        cif_path = (\n            output_path\n            / target_id\n            / target_id\n            / f\"seed_{seed}\"\n            / \"predictions\"\n            / f\"{target_id}_seed_{seed}_sample_{i}.cif\"\n        )\n\n        if cif_path.exists():\n            pred_df = extract_c1_atoms(cif_path)\n            coords = np.zeros((len(sequence), 3))\n            n_atoms = min(len(pred_df), len(sequence))\n            coords[:n_atoms] = pred_df[[\"x\", \"y\", \"z\"]].values[:n_atoms]\n            predictions.append(coords)\n\n    return predictions\n\n\n@contextlib.contextmanager\ndef protenix_context(protenix_dir: str):\n    original_dir = os.getcwd()\n    os.chdir(protenix_dir)\n    try:\n        yield\n    finally:\n        os.chdir(original_dir)\n\n\n# ===========================================================================\n# Combined prediction pipeline (Template + Protenix + de novo fill)\n# ===========================================================================\ndef generate_predictions_batch(\n    sequences_df: pd.DataFrame,\n    train_seqs_df: pd.DataFrame,\n    train_coords_dict: dict,\n    dataset_name: str,\n    aligner,\n    segments_map: dict,\n    is_kaggle: bool,\n    data_path: str,\n    output_path: str,\n    protenix_dir: str,\n    use_temporal_cutoff: bool,\n    use_protenix: bool,\n    min_similarity: float,\n    min_percent_identity: float,\n    show_alignment_details: bool,\n):\n    start_time = time.time()\n    total_targets = len(sequences_df)\n\n    print(f\"\\n{'='*70}\")\n    print(f\"Predicting {total_targets} {dataset_name} sequences\")\n    print(f\"{'='*70}\")\n\n    # Metadata (kept local to this run)\n    template_info_dict = {}\n    prediction_metadata_dict = {}\n\n    def record_template_info(target_id, template_id, similarity, percent_identity):\n        if target_id not in template_info_dict:\n            template_info_dict[target_id] = {\"template_ids\": [], \"similarities\": [], \"percent_identities\": []}\n        template_info_dict[target_id][\"template_ids\"].append(template_id)\n        template_info_dict[target_id][\"similarities\"].append(similarity)\n        template_info_dict[target_id][\"percent_identities\"].append(percent_identity)\n\n    def record_prediction_metadata(target_id, pred_num, source, template_id=None, similarity=None, percent_identity=None):\n        if target_id not in prediction_metadata_dict:\n            prediction_metadata_dict[target_id] = {}\n        prediction_metadata_dict[target_id][pred_num] = {\n            \"source\": source,\n            \"template_id\": template_id if source == \"template\" else None,\n            \"similarity\": similarity if source == \"template\" else None,\n            \"percent_identity\": percent_identity if source == \"template\" else None,\n        }\n\n    # ---- Phase 1: Template predictions ----\n    print(\"\\nPHASE 1: Template-based predictions\")\n\n    template_predictions = {}\n    protenix_queue = {}\n\n    for _, row in sequences_df.iterrows():\n        target_id = row[\"target_id\"]\n        sequence = row[\"sequence\"]\n\n        temporal_cutoff = None\n        if use_temporal_cutoff and \"temporal_cutoff\" in row:\n            temporal_cutoff = row.get(\"temporal_cutoff\", None)\n\n        print(\"\\n\" + \"-\" * 70)\n        print(f\"Target: {target_id} ({len(sequence)} nt)\")\n\n        preds = []\n        pred_num = 1\n\n        similar_seqs = find_similar_sequences_detailed(\n            query_seq=sequence,\n            train_seqs_df=train_seqs_df,\n            train_coords_dict=train_coords_dict,\n            aligner=aligner,\n            temporal_cutoff=temporal_cutoff,\n            top_n=5,\n        )\n\n        if similar_seqs:\n            for i, (tmpl_id, tmpl_seq, similarity, tmpl_coords, pct_id, aligned_q, aligned_t) in enumerate(similar_seqs):\n                if (similarity < min_similarity or pct_id < min_percent_identity) and len(tmpl_seq) < 500:\n                    if show_alignment_details:\n                        print(\n                            f\"  Template {i+1}: {tmpl_id} SKIPPED (sim={similarity:.3f}, id={pct_id:.1f}%) - below threshold\"\n                        )\n                    break\n\n                if use_protenix and len(aligned_q) < 100 and i == 4:\n                    print(f\"  Sequence length {len(aligned_q)} short, leaving 1 space for Protenix\")\n                    break\n\n                record_template_info(target_id, tmpl_id, similarity, pct_id)\n                record_prediction_metadata(target_id, pred_num, \"template\", tmpl_id, similarity, pct_id)\n\n                if show_alignment_details:\n                    print(f\"  Template {i+1}: {tmpl_id} (sim={similarity:.3f}, id={pct_id:.1f}%)\")\n\n                adapted = adapt_template_to_query(sequence, tmpl_seq, tmpl_coords, aligner=aligner)\n                refined = adaptive_rna_constraints(adapted, target_id, segments_map, confidence=similarity)\n                preds.append(refined)\n\n                pred_num += 1\n                if len(preds) >= 5:\n                    break\n\n        template_predictions[target_id] = preds\n        n_needed = 5 - len(preds)\n\n        if n_needed > 0:\n            print(f\"  -> {len(preds)} from templates, {n_needed} slots for Protenix\")\n            protenix_queue[target_id] = (n_needed, pred_num, sequence)\n        else:\n            print(\"  -> All 5 predictions from templates\")\n\n    template_time = time.time() - start_time\n    print(f\"\\nPhase 1 done: {template_time:.1f}s | {len(protenix_queue)} targets need Protenix\")\n\n    # ---- Phase 2: Protenix predictions ----\n    protenix_predictions = {}\n\n    if protenix_queue and use_protenix:\n        print(f\"\\nPHASE 2: Protenix for {len(protenix_queue)} targets\")\n\n        protenix_output_path = Path(output_path) / f\"{dataset_name}_protenix\"\n        protenix_output_path.mkdir(parents=True, exist_ok=True)\n        input_path = Path(data_path)\n\n        for i, (target_id, (n_needed, _next_pred, sequence)) in enumerate(protenix_queue.items()):\n            print(f\"\\n  [{i+1}/{len(protenix_queue)}] {target_id} ({len(sequence)} nt, need {n_needed})\")\n\n            try:\n                prepare_protenix_json(target_id, sequence, protenix_output_path, input_path)\n\n                t0 = time.time()\n                run_protenix_inference(\n                    is_kaggle=is_kaggle,\n                    data_path=data_path.rstrip(\"/\"),\n                    target_id=target_id,\n                    sequence=sequence,\n                    output_path=protenix_output_path,\n                    input_path=input_path,\n                    protenix_dir=protenix_dir,\n                    seed=101,\n                    n_cycle=10,\n                    n_sample=n_needed,\n                    n_step=200,\n                )\n                print(f\"    Done in {(time.time()-t0)/60:.1f} min\")\n\n                preds = get_protenix_predictions(target_id, sequence, protenix_output_path, seed=101, n_sample=n_needed)\n                protenix_predictions[target_id] = preds\n                print(f\"    Got {len(preds)} Protenix predictions\")\n            except Exception as e:\n                print(f\"    Protenix FAILED: {e}\")\n                print(\"    Will fall back to de novo\")\n                protenix_predictions[target_id] = None\n\n    elif protenix_queue and not use_protenix:\n        print(f\"\\nPHASE 2: Protenix disabled, will use de novo for {len(protenix_queue)} targets\")\n\n    # ---- Phase 3: Combine and format ----\n    print(\"\\nPHASE 3: Combining predictions\")\n\n    all_rows = []\n\n    for _, row in sequences_df.iterrows():\n        target_id = row[\"target_id\"]\n        sequence = row[\"sequence\"]\n\n        predictions = list(template_predictions[target_id])\n        pred_num = len(predictions) + 1\n\n        # Fill with Protenix predictions\n        if target_id in protenix_queue:\n            ptx_preds = protenix_predictions.get(target_id)\n            if ptx_preds:\n                for coords in ptx_preds:\n                    record_prediction_metadata(target_id, pred_num, \"protenix\")\n                    predictions.append(coords)\n                    pred_num += 1\n                    if len(predictions) >= 5:\n                        break\n\n        # Fill remaining with de novo\n        while len(predictions) < 5:\n            record_prediction_metadata(target_id, pred_num, \"de_novo\")\n            seed_val = hash(target_id) % 10000 + len(predictions) * 1000\n            de_novo = generate_rna_structure(sequence, seed=seed_val)\n            refined = adaptive_rna_constraints(de_novo, target_id, segments_map, confidence=0.2)\n            predictions.append(refined)\n            pred_num += 1\n\n        # Output rows\n        for j in range(len(sequence)):\n            pred_row = {\"ID\": f\"{target_id}_{j+1}\", \"resname\": sequence[j], \"resid\": j + 1}\n            for i_pred in range(5):\n                pred_row[f\"x_{i_pred+1}\"] = predictions[i_pred][j][0]\n                pred_row[f\"y_{i_pred+1}\"] = predictions[i_pred][j][1]\n                pred_row[f\"z_{i_pred+1}\"] = predictions[i_pred][j][2]\n            all_rows.append(pred_row)\n\n    submission_df = pd.DataFrame(all_rows)\n    column_order = [\"ID\", \"resname\", \"resid\"]\n    for i in range(1, 6):\n        for coord in [\"x\", \"y\", \"z\"]:\n            column_order.append(f\"{coord}_{i}\")\n    submission_df = submission_df[column_order]\n\n    total_time = time.time() - start_time\n    n_template_only = sum(1 for tid in sequences_df[\"target_id\"] if tid not in protenix_queue)\n\n    print(f\"\\n{'='*70}\")\n    print(f\"{dataset_name.upper()} PREDICTIONS COMPLETE\")\n    print(f\"  Template-only targets: {n_template_only}\")\n    print(f\"  Targets with Protenix: {len(protenix_queue)}\")\n    print(f\"  Total residues: {len(submission_df)}\")\n    print(f\"  Runtime: {total_time:.1f}s ({total_time/60:.1f} min)\")\n    print(f\"{'='*70}\\n\")\n\n    return submission_df, template_info_dict, prediction_metadata_dict\n\n\n# ===========================================================================\n# Validation scoring (USalign) — kept from notebook, minimal changes\n# ===========================================================================\ndef parse_tmscore_output(output: str):\n    tm_matches = re.findall(r\"TM-score=\\s+([\\d.]+)\", output)\n    if len(tm_matches) < 2:\n        raise ValueError(\"No TM score found in USalign output\")\n\n    tm_score = float(tm_matches[1])\n\n    rmsd_match = re.search(r\"RMSD=\\s*([\\d.]+)\", output, re.IGNORECASE)\n    rmsd = float(rmsd_match.group(1)) if rmsd_match else None\n\n    return tm_score, rmsd\n\n\ndef sanitize(xyz):\n    min_coord = -999.999\n    max_coord = 9999.999\n    return min(max(xyz, min_coord), max_coord)\n\n\ndef write_target_line(\n    atom_name,\n    atom_serial,\n    residue_name,\n    chain_id,\n    residue_num,\n    x_coord,\n    y_coord,\n    z_coord,\n    occupancy=1.0,\n    b_factor=0.0,\n    atom_type=\"P\",\n) -> str:\n    return (\n        f\"ATOM  {atom_serial:>5d}  {atom_name:4s}{residue_name:>3s} {chain_id:1s}{residue_num:>4d}    \"\n        f\"{sanitize(x_coord):>8.3f}{sanitize(y_coord):>8.3f}{sanitize(z_coord):>8.3f}\"\n        f\"{occupancy:>6.2f}{b_factor:>6.2f}           {atom_type}\\n\"\n    )\n\n\ndef write2pdb(df: pd.DataFrame, xyz_id: int, target_path: str) -> int:\n    resolved_cnt = 0\n    with open(target_path, \"w\") as fh:\n        for _, row in df.iterrows():\n            x = row[f\"x_{xyz_id}\"]\n            y = row[f\"y_{xyz_id}\"]\n            z = row[f\"z_{xyz_id}\"]\n            if x > -1e6 and y > -1e6 and z > -1e6:\n                resolved_cnt += 1\n                resid_num = int(row[\"resid\"])\n                fh.write(\n                    write_target_line(\"C1'\", resid_num, row[\"resname\"], \"A\", resid_num, x, y, z, atom_type=\"C\")\n                )\n    return resolved_cnt\n\n\ndef run_usalign_raw(\n    predicted_pdb: str,\n    native_pdb: str,\n    usalign_bin=\"USalign\",\n    align_sequence=False,\n    tmscore=None,\n    show_output: bool = False,\n) -> str:\n    cmd = f'{usalign_bin} {predicted_pdb} {native_pdb} -atom \" C1\\'\"'\n    if tmscore is not None:\n        cmd += f\" -TMscore {tmscore}\"\n        if int(tmscore) == 0:\n            cmd += \" -mm 1 -ter 0\"\n    elif not align_sequence:\n        cmd += \" -TMscore 1\"\n\n    if show_output:\n        print(\"-\" * 100)\n        print(f\"Running {cmd}\")\n\n    proc = subprocess.run(cmd, shell=True, capture_output=True, text=True)\n    res = proc.stdout + proc.stderr\n\n    if show_output:\n        print(res)\n\n    return res\n\n\ndef score_simple(solution: pd.DataFrame, submission: pd.DataFrame, usalign_bin: str, show_output: bool = False):\n    sol = solution.copy()\n    sub = submission.copy()\n\n    sol[\"target_id\"] = sol[\"ID\"].apply(lambda x: \"_\".join(str(x).split(\"_\")[:-1]))\n    sub[\"target_id\"] = sub[\"ID\"].apply(lambda x: \"_\".join(str(x).split(\"_\")[:-1]))\n\n    results = []\n    scores_df = pd.DataFrame()\n\n    for target_id, group_native in sol.groupby(\"target_id\"):\n        group_predicted = sub[sub[\"target_id\"] == target_id]\n\n        native_with_coords = []\n        for native_cnt in range(1, 41):\n            native_pdb = f\"native_{target_id}_{native_cnt}.pdb\"\n            resolved_native = write2pdb(group_native, native_cnt, native_pdb)\n            if resolved_native > 0:\n                native_with_coords.append(native_cnt)\n            else:\n                if os.path.exists(native_pdb):\n                    os.remove(native_pdb)\n\n        if not native_with_coords:\n            raise ValueError(f\"No native models with coordinates for target {target_id}\")\n\n        best_per_pred = []\n        for pred_cnt in range(1, 6):\n            predicted_pdb = f\"predicted_{target_id}_{pred_cnt}.pdb\"\n            resolved_pred = write2pdb(group_predicted, pred_cnt, predicted_pdb)\n            if resolved_pred <= 2:\n                best_per_pred.append(0.0)\n                continue\n\n            scores = []\n            for native_cnt in native_with_coords:\n                native_pdb = f\"native_{target_id}_{native_cnt}.pdb\"\n                out = run_usalign_raw(\n                    predicted_pdb,\n                    native_pdb,\n                    usalign_bin=usalign_bin,\n                    align_sequence=False,\n                    tmscore=1,\n                    show_output=show_output,\n                )\n                s, r = parse_tmscore_output(out)\n                scores.append(s)\n\n                idx = len(scores_df)\n                scores_df.loc[idx, \"target_id\"] = str(target_id)\n                scores_df.loc[idx, \"native_cnt\"] = int(native_cnt)\n                scores_df.loc[idx, \"pred_cnt\"] = int(pred_cnt)\n                scores_df.loc[idx, \"tm_score\"] = float(s)\n                scores_df.loc[idx, \"rmsd\"] = float(r) if r is not None else np.nan\n\n            best_per_pred.append(max(scores))\n\n        results.append(max(best_per_pred))\n\n    score_mean = float(sum(results) / len(results)) if len(results) > 0 else 0.0\n    return score_mean, scores_df\n\n\n# ===========================================================================\n# Main\n# ===========================================================================\ndef main():\n    # Seed\n    seed_everything(42)\n\n    # Detect environment + paths\n    is_kaggle, data_path = detect_environment(CONFIG[\"DATA_PATH\"])\n    paths = build_paths(is_kaggle, data_path)\n\n    print(f\"Data path: {paths['DATA_PATH']}\")\n    print(f\"Protenix dir: {paths['PROTENIX_DIR']}\")\n\n    # Kaggle setup\n    if is_kaggle:\n        kaggle_setup(paths)\n    else:\n        # Local helpers path (as your notebook did)\n        sys.path.insert(0, os.path.join(os.getcwd(), \"kaggle_util_ds\"))\n\n    # Load data\n    train_seqs, validation_seqs, test_seqs, train_labels, validation_labels = load_sequence_and_label_data(\n        paths[\"DATA_PATH\"]\n    )\n\n    # Process coords\n    print(\"Processing coordinates...\")\n    train_coords_dict = process_labels(train_labels)\n\n    combined_seqs = pd.concat([train_seqs, validation_seqs], ignore_index=True)\n    combined_labels = pd.concat([train_labels, validation_labels], ignore_index=True)\n    combined_coords_dict = process_labels(combined_labels)\n\n    print(f\"Processed {len(train_coords_dict)} training structures\")\n    print(f\"Processed {len(combined_coords_dict)} combined (train+val) structures\")\n\n    # Segment maps\n    validation_segments_map, _ = build_segments_map(validation_seqs)\n    test_segments_map, _ = build_segments_map(test_seqs)\n\n    print(f\"Built segment maps: {len(validation_segments_map)} validation, {len(test_segments_map)} test\")\n\n    # Aligner\n    aligner = make_aligner()\n\n    # Debug: optionally restrict validation to shortest\n    if CONFIG[\"DEBUG\"] and CONFIG[\"SHOW_VALIDATION\"]:\n        shortest_ids = (\n            validation_seqs.assign(seq_len=validation_seqs[\"sequence\"].str.len())\n            .nsmallest(5, \"seq_len\")[\"target_id\"]\n            .tolist()\n        )\n        validation_seqs = validation_seqs[validation_seqs[\"target_id\"].isin(shortest_ids)].reset_index(drop=True)\n        validation_labels = validation_labels[\n            validation_labels[\"ID\"].str.rsplit(\"_\", n=1).str[0].isin(shortest_ids)\n        ].reset_index(drop=True)\n        print(f\"DEBUG: Using {len(validation_seqs)} shortest sequences for validation\")\n\n    # Validation run (templates from TRAIN only, no leakage)\n    if CONFIG[\"SHOW_VALIDATION\"]:\n        validation_predictions, _tmpl_info, _pred_meta = generate_predictions_batch(\n            sequences_df=validation_seqs,\n            train_seqs_df=train_seqs,\n            train_coords_dict=train_coords_dict,\n            dataset_name=\"validation\",\n            aligner=aligner,\n            segments_map=validation_segments_map,\n            is_kaggle=is_kaggle,\n            data_path=paths[\"DATA_PATH\"],\n            output_path=paths[\"OUTPUT_PATH\"],\n            protenix_dir=paths[\"PROTENIX_DIR\"],\n            use_temporal_cutoff=True,\n            use_protenix=bool(CONFIG[\"USE_PROTENIX\"]),\n            min_similarity=float(CONFIG[\"MIN_SIMILARITY\"]),\n            min_percent_identity=float(CONFIG[\"MIN_PERCENT_IDENTITY\"]),\n            show_alignment_details=bool(CONFIG[\"SHOW_ALIGNMENT_DETAILS\"]),\n        )\n\n        validation_predictions.to_csv(\"validation_predictions.csv\", index=False)\n        print(\"Saved: validation_predictions.csv\")\n\n        # Scoring (single-chain version)\n        mean_tm_score, scores_df = score_simple(\n            solution=validation_labels,\n            submission=validation_predictions,\n            usalign_bin=paths[\"USALIGN_BIN\"],\n            show_output=False,\n        )\n        scores_df.to_csv(\"validation_detailed_scores.csv\", index=False, float_format=\"%.3f\")\n\n        print(f\"\\nValidation mean TM-score (best-of-5, simple scorer): {mean_tm_score:.4f}\")\n        print(\"Saved: validation_detailed_scores.csv\")\n\n    # Test run (templates from TRAIN+VAL, Protenix fill)\n    if CONFIG[\"MAKE_SUBMISSION\"]:\n        test_predictions, _tmpl_info, _pred_meta = generate_predictions_batch(\n            sequences_df=test_seqs,\n            train_seqs_df=combined_seqs,\n            train_coords_dict=combined_coords_dict,\n            dataset_name=\"test\",\n            aligner=aligner,\n            segments_map=test_segments_map,\n            is_kaggle=is_kaggle,\n            data_path=paths[\"DATA_PATH\"],\n            output_path=paths[\"OUTPUT_PATH\"],\n            protenix_dir=paths[\"PROTENIX_DIR\"],\n            use_temporal_cutoff=False,\n            use_protenix=bool(CONFIG[\"USE_PROTENIX\"]),\n            min_similarity=float(CONFIG[\"MIN_SIMILARITY\"]),\n            min_percent_identity=float(CONFIG[\"MIN_PERCENT_IDENTITY\"]),\n            show_alignment_details=bool(CONFIG[\"SHOW_ALIGNMENT_DETAILS\"]),\n        )\n\n        test_predictions.to_csv(\"submission.csv\", index=False)\n        print(\"Saved: submission.csv\")\n\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T03:39:55.065396Z","iopub.execute_input":"2026-02-08T03:39:55.065732Z","iopub.status.idle":"2026-02-08T03:43:25.947347Z","shell.execute_reply.started":"2026-02-08T03:39:55.065703Z","shell.execute_reply":"2026-02-08T03:43:25.946312Z"}},"outputs":[],"execution_count":null}]}