{"metadata":{"kernelspec":{"display_name":"protenixpy310","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.19"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":118765,"databundleVersionId":15231210,"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":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport sys\n\nIS_KAGGLE = True\nDATA_PATH = '/kaggle/input/stanford-rna-3d-folding-2/'\nOUTPUT_PATH = '/kaggle/working/output'\nUSALIGN_BIN = '/kaggle/working/USalign'\nPROTENIX_DIR = '/kaggle/working/Protenix'\n! cp /kaggle/input/protenix-packages/packages/USalign /kaggle/working/\n! chmod +x /kaggle/working/USalign\nsys.path.insert(0, '/kaggle/input/rna-3d-utils/')\n\nprint(f\"Data path: {DATA_PATH}\")\nprint(f\"Protenix dir: {PROTENIX_DIR}\")","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if IS_KAGGLE:\n    !cp -r /kaggle/input/protenix-packages/packages /kaggle/working\n    %cd /kaggle/working/packages\n    !pip install --no-deps --exists-action=i *.whl\n    %cd /kaggle/working\n\n    !mv /kaggle/working/packages/ihm-2.3/ihm-2.3 /kaggle/working\n    !mv /kaggle/working/packages/modelcif-0.7/modelcif-0.7 /kaggle/working\n\n    !pip install /kaggle/working/ihm-2.3\n    !pip install /kaggle/working/modelcif-0.7\n\n    !rm -rf /kaggle/working/ihm-2.3\n    !rm -rf /kaggle/working/modelcif-0.7\n\n    !pip install /kaggle/input/biopython/biopython-1.85-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n    !pip install /kaggle/input/ml-collections/ml_collections-1.0.0-py3-none-any.whl\n\n    !rm -rf /kaggle/working/packages\n\n    !cp -r /kaggle/input/protenix-mg-packages/protenix_mg_packages /kaggle/working\n    %cd /kaggle/working/protenix_mg_packages\n    !pip install --no-deps --exists-action=i *.whl\n    %cd /kaggle/working\n    !rm -rf /kaggle/working/protenix_mg_packages\n\n    !cp -R /kaggle/input/protenix-rmsa-repo/protenix_kaggle /kaggle/working/\n    !mv protenix_kaggle Protenix","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json\nimport numpy as np\nimport pandas as pd\nimport time\nimport random\nimport warnings\nimport contextlib\nfrom pathlib import Path\nfrom Bio.Align import PairwiseAligner\n\nwarnings.filterwarnings('ignore')\n\ndef seed_everything(seed: int = 42):\n    random.seed(seed)\n    np.random.seed(seed)\n\nseed_everything(42)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_seqs = pd.read_csv(DATA_PATH + 'train_sequences.csv')\nvalidation_seqs = pd.read_csv(DATA_PATH + 'validation_sequences.csv')\ntest_seqs = pd.read_csv(DATA_PATH + 'test_sequences.csv')\ntrain_labels = pd.read_csv(DATA_PATH + 'train_labels.csv')\nvalidation_labels = pd.read_csv(DATA_PATH + 'validation_labels.csv')\n\nSHOW_VALIDATION = False\nMAKE_SUBMISSION = True\nUSE_PROTENIX = True\n\nMIN_SIMILARITY = 0.45\nMIN_PERCENT_IDENTITY = 55\nTOP_TEMPLATES = 10\nENSEMBLE_WEIGHTS = [0.40, 0.25, 0.15, 0.12, 0.08]","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def make_aligner():\n    al = PairwiseAligner()\n    al.mode = 'global'\n    al.match_score = 2.5\n    al.mismatch_score = -2.0\n    al.open_gap_score = -10\n    al.extend_gap_score = -0.5\n    al.query_left_open_gap_score = -10\n    al.query_left_extend_gap_score = -0.5\n    al.query_right_open_gap_score = -10\n    al.query_right_extend_gap_score = -0.5\n    al.target_left_open_gap_score = -10\n    al.target_left_extend_gap_score = -0.5\n    al.target_right_open_gap_score = -10\n    al.target_right_extend_gap_score = -0.5\n    return al\n\n_aligner = make_aligner()\n\ndef parse_fasta(fasta_content: str):\n    out = {}\n    cur = None\n    seq_parts = []\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(seq_parts)\n            cur = line[1:].split()[0]\n            seq_parts = []\n        else:\n            seq_parts.append(line.replace(\" \", \"\"))\n    if cur is not None:\n        out[cur] = \"\".join(seq_parts)\n    return out\n\ndef parse_stoichiometry(stoich: str):\n    if pd.isna(stoich) or str(stoich).strip() == \"\":\n        return []\n    out = []\n    for part in str(stoich).split(';'):\n        ch, cnt = part.split(':')\n        out.append((ch.strip(), int(cnt)))\n    return out\n\ndef get_chain_segments(row):\n    seq = row['sequence']\n    stoich = row.get('stoichiometry', '')\n    all_seq = row.get('all_sequences', '')\n    \n    if pd.isna(stoich) or pd.isna(all_seq) or str(stoich).strip() == \"\" or str(all_seq).strip() == \"\":\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        for ch, cnt in order:\n            base = chain_dict.get(ch)\n            if base is None:\n                return [(0, len(seq))]\n            for _ in range(cnt):\n                L = len(base)\n                segs.append((pos, pos + L))\n                pos += L\n        if pos != len(seq):\n            return [(0, len(seq))]\n        return segs\n    except Exception:\n        return [(0, len(seq))]\n\ndef build_segments_map(df):\n    seg_map = {}\n    stoich_map = {}\n    for _, r in df.iterrows():\n        tid = r['target_id']\n        seg_map[tid] = get_chain_segments(r)\n        stoich_map[tid] = str(r.get('stoichiometry', '') if not pd.isna(r.get('stoichiometry', '')) else '')\n    return seg_map, stoich_map\n\ndef process_labels(labels_df):\n    coords_dict = {}\n    prefixes = labels_df['ID'].str.rsplit('_', n=1).str[0]\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    return coords_dict\n\ndef compute_sequence_features(seq):\n    length = len(seq)\n    gc_content = (seq.count('G') + seq.count('C')) / length if length > 0 else 0\n    au_content = (seq.count('A') + seq.count('U')) / length if length > 0 else 0\n    return {\n        'length': length,\n        'gc_content': gc_content,\n        'au_content': au_content,\n        'complexity': len(set(seq)) / 4.0\n    }\n\ndef enhanced_template_selection(query_seq, train_seqs_df, train_coords_dict, temporal_cutoff=None, top_n=10):\n    similar_seqs = []\n    query_features = compute_sequence_features(query_seq)\n    \n    if temporal_cutoff is not None:\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, train_seq = row['target_id'], row['sequence']\n        if target_id not in train_coords_dict:\n            continue\n        \n        len_ratio = abs(len(train_seq) - len(query_seq)) / max(len(train_seq), len(query_seq))\n        if len_ratio > 0.4:\n            continue\n        \n        template_features = compute_sequence_features(train_seq)\n        feature_similarity = 1.0 - abs(query_features['gc_content'] - template_features['gc_content'])\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        percent_identity = 100 * identical / len(query_seq)\n        \n        combined_score = 0.7 * normalized_score + 0.2 * (percent_identity / 100) + 0.1 * feature_similarity\n        \n        aligned_query, aligned_template = _build_aligned_strings(query_seq, train_seq, alignment)\n        \n        similar_seqs.append((\n            target_id, train_seq, combined_score, normalized_score,\n            train_coords_dict[target_id], percent_identity,\n            aligned_query, aligned_template\n        ))\n    \n    similar_seqs.sort(key=lambda x: x[2], reverse=True)\n    return similar_seqs[:top_n]\n\ndef _build_aligned_strings(query_seq, template_seq, 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        while ti < ts:\n            aligned_q.append('-')\n            aligned_t.append(template_seq[ti])\n            ti += 1\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        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    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\ndef adapt_template_to_query(query_seq, template_seq, template_coords):\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            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.8, 0, 0]\n            elif next_v >= 0:\n                new_coords[i] = new_coords[next_v] + [3.8, 0, 0]\n            else:\n                new_coords[i] = [i * 3.8, 0, 0]\n    \n    return np.nan_to_num(new_coords)\n\ndef enhanced_rna_constraints(coordinates, target_id, segments_map, confidence=1.0, passes=3):\n    coords = coordinates.copy()\n    segments = segments_map.get(target_id, [(0, len(coords))])\n    \n    base_strength = 0.85 * (1.0 - min(confidence, 0.98))\n    base_strength = max(base_strength, 0.01)\n    \n    for pass_idx in range(passes):\n        strength = base_strength * (1.0 - pass_idx * 0.15)\n        \n        for (s, e) in segments:\n            X = coords[s:e]\n            L = e - s\n            if L < 3:\n                coords[s:e] = X\n                continue\n            \n            d = X[1:] - X[:-1]\n            dist = np.linalg.norm(d, axis=1) + 1e-6\n            target_bond = 5.9\n            scale = (target_bond - dist) / dist\n            adj = (d * scale[:, None]) * (0.28 * strength)\n            X[:-1] -= adj\n            X[1:] += adj\n            \n            if L >= 3:\n                d2 = X[2:] - X[:-2]\n                dist2 = np.linalg.norm(d2, axis=1) + 1e-6\n                target2 = 10.4\n                scale2 = (target2 - dist2) / dist2\n                adj2 = (d2 * scale2[:, None]) * (0.15 * strength)\n                X[:-2] -= adj2\n                X[2:] += adj2\n            \n            if L >= 4:\n                lap = 0.5 * (X[:-2] + X[2:]) - X[1:-1]\n                X[1:-1] += (0.08 * strength) * lap\n            \n            if L >= 6:\n                backbone_smooth = 0.33 * (X[:-2] + X[1:-1] + X[2:])\n                X[1:-1] = (1 - 0.12 * strength) * X[1:-1] + (0.12 * strength) * backbone_smooth\n            \n            if L >= 30:\n                k = min(L, 180) if L > 250 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.5)\n                if np.any(mask):\n                    force = (3.5 - distm) / distm\n                    vec = (diff * force[:, :, None] * mask[:, :, None]).sum(axis=1)\n                    X[idx] += (0.018 * strength) * vec\n            \n            coords[s:e] = X\n    \n    return coords\n\ndef _rotmat(axis, ang):\n    axis = np.asarray(axis, float)\n    axis = axis / (np.linalg.norm(axis) + 1e-12)\n    x, y, z = axis\n    c, s = np.cos(ang), np.sin(ang)\n    C = 1.0 - c\n    return np.array([\n        [c + x*x*C,     x*y*C - z*s, x*z*C + y*s],\n        [y*x*C + z*s,   c + y*y*C,   y*z*C - x*s],\n        [z*x*C - y*s,   z*y*C + x*s, c + z*z*C]\n    ], dtype=float)\n\ndef apply_hinge(coords, seg, rng, max_angle_deg=20):\n    s, e = seg\n    L = e - s\n    if L < 35:\n        return coords\n    pivot = s + int(rng.integers(12, L - 12))\n    axis = rng.normal(size=3)\n    ang = np.deg2rad(float(rng.uniform(-max_angle_deg, max_angle_deg)))\n    R = _rotmat(axis, ang)\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\ndef weighted_ensemble_prediction(templates, query_seq, segments_map, target_id):\n    if not templates:\n        return None\n    \n    predictions = []\n    weights = ENSEMBLE_WEIGHTS[:len(templates)]\n    weights = np.array(weights)\n    weights = weights / weights.sum()\n    \n    for (tmpl_id, tmpl_seq, combined_score, sim, tmpl_coords, pct_id, _, _), weight in zip(templates, weights):\n        adapted = adapt_template_to_query(query_seq, tmpl_seq, tmpl_coords)\n        refined = enhanced_rna_constraints(adapted, target_id, segments_map, confidence=sim, passes=3)\n        predictions.append((refined, weight))\n    \n    if len(predictions) == 1:\n        return predictions[0][0]\n    \n    ensemble_coords = np.zeros_like(predictions[0][0])\n    for coords, weight in predictions:\n        ensemble_coords += weight * coords\n    \n    final = enhanced_rna_constraints(ensemble_coords, target_id, segments_map, confidence=0.85, passes=2)\n    \n    return final\n\ndef generate_rna_structure(sequence, seed=None):\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        angle = i * 0.58\n        coords[i] = [11.0 * np.cos(angle), 11.0 * np.sin(angle), i * 2.6]\n    return coords\n\ntrain_coords_dict = process_labels(train_labels)\ncombined_seqs = pd.concat([train_seqs, validation_seqs], ignore_index=True)\ncombined_labels = pd.concat([train_labels, validation_labels], ignore_index=True)\ncombined_coords_dict = process_labels(combined_labels)\n\nvalidation_segments_map, _ = build_segments_map(validation_seqs)\ntest_segments_map, _ = build_segments_map(test_seqs)\n\nfrom biotite.structure.io.pdbx import CIFFile, get_structure\nimport contextlib\n\ndef extract_c1_atoms(cif_path):\n    cif_file = CIFFile.read(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    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    return df[[\"res_name\", \"res_id\", \"x\", \"y\", \"z\"]]\n\ndef prepare_protenix_json(target_id, sequence, output_path, input_path, max_length=400):\n    if len(sequence) <= max_length:\n        input_json = [{\n            \"sequences\": [{\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            \"name\": target_id,\n        }]\n    else:\n        input_json = [{\n            \"sequences\": [{\n                \"rnaSequence\": {\n                    \"sequence\": sequence[:max_length],\n                    \"count\": 1,\n                }\n            }],\n            \"name\": target_id,\n        }]\n    \n    json_path = Path(output_path) / \"input_json\" / f\"{target_id}.json\"\n    json_path.parent.mkdir(parents=True, exist_ok=True)\n    with open(json_path, \"w\") as f:\n        json.dump(input_json, f, indent=4)\n\ndef run_protenix_inference(target_id, sequence, output_path, input_path,\n                           seed=101, n_cycle=12, n_sample=5, n_step=250, max_length=400):\n    if IS_KAGGLE:\n        checkpoint_path = \"/kaggle/input/protenix-finetuned-rna3db-all-1599/1599_ema_0.999.pt\"\n    else:\n        checkpoint_path = f\"{DATA_PATH}/protenix_chpt/1599_ema_0.999.pt\"\n    \n    output_path = Path(output_path)\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    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    from runner.inference import run\n    run()\n\ndef get_protenix_predictions(target_id, sequence, output_path, seed=101, n_sample=5):\n    output_path = Path(output_path)\n    predictions = []\n    for i in range(n_sample):\n        cif_path = (output_path / target_id / target_id / f\"seed_{seed}\" /\n                    \"predictions\" / f\"{target_id}_seed_{seed}_sample_{i}.cif\")\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    return predictions\n\n@contextlib.contextmanager\ndef protenix_context():\n    original_dir = os.getcwd()\n    os.chdir(PROTENIX_DIR)\n    try:\n        yield\n    finally:\n        os.chdir(original_dir)\n\ntemplate_info_dict = {}\nprediction_metadata_dict = {}\n\ndef 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] = {\n            'template_ids': [], 'similarities': [], 'percent_identities': []\n        }\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\ndef record_prediction_metadata(target_id, pred_num, source, template_id=None,\n                               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\ndef predict_with_enhanced_templates(sequence, target_id, train_seqs_df, train_coords_dict,\n                                   segments_map, n_predictions=5, temporal_cutoff=None):\n    predictions = []\n    pred_num = 1\n    \n    print(f\"\\nTarget: {target_id} ({len(sequence)} nt)\")\n    \n    similar_seqs = enhanced_template_selection(\n        sequence, train_seqs_df, train_coords_dict,\n        temporal_cutoff=temporal_cutoff, top_n=TOP_TEMPLATES\n    )\n    \n    if similar_seqs:\n        top_templates = []\n        for i, (tmpl_id, tmpl_seq, combined_score, similarity, tmpl_coords,\n                pct_id, aligned_q, aligned_t) in enumerate(similar_seqs):\n            \n            if (similarity < MIN_SIMILARITY or pct_id < MIN_PERCENT_IDENTITY) and len(tmpl_seq) < 500:\n                print(f\"  Template {i+1}: {tmpl_id} SKIPPED (sim={similarity:.3f}, id={pct_id:.1f}%)\")\n                break\n            \n            if USE_PROTENIX and len(aligned_q) < 100 and i == 4:\n                print(f\"  Leaving 1 slot for Protenix\")\n                break\n            \n            record_template_info(target_id, tmpl_id, similarity, pct_id)\n            print(f\"  Template {i+1}: {tmpl_id} (sim={similarity:.3f}, id={pct_id:.1f}%)\")\n            \n            top_templates.append((tmpl_id, tmpl_seq, combined_score, similarity, \n                                 tmpl_coords, pct_id, aligned_q, aligned_t))\n            \n            if len(top_templates) >= min(5, n_predictions):\n                break\n        \n        if top_templates:\n            ensemble_pred = weighted_ensemble_prediction(top_templates, sequence, \n                                                        segments_map, target_id)\n            if ensemble_pred is not None:\n                record_prediction_metadata(target_id, pred_num, 'ensemble_template',\n                                         top_templates[0][0], top_templates[0][3], \n                                         top_templates[0][5])\n                predictions.append(ensemble_pred)\n                pred_num += 1\n    \n    n_from_templates = len(predictions)\n    n_needed = n_predictions - n_from_templates\n    \n    if n_needed > 0:\n        print(f\"  -> {n_from_templates} ensemble pred, {n_needed} slots for Protenix\")\n    else:\n        print(f\"  -> {n_predictions} predictions from ensemble\")\n    \n    return predictions, n_needed, pred_num\n\ndef generate_predictions_batch(sequences_df, train_seqs_df, train_coords_dict,\n                               dataset_name, use_temporal_cutoff=True,\n                               protenix_output_path=None):\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    segments_map, _ = build_segments_map(sequences_df)\n    \n    print(\"\\nPHASE 1: Enhanced template-based ensemble 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        temporal_cutoff = row.get('temporal_cutoff', None) if use_temporal_cutoff else None\n        \n        preds, n_needed, next_pred = predict_with_enhanced_templates(\n            sequence, target_id, train_seqs_df, train_coords_dict,\n            segments_map, n_predictions=5, temporal_cutoff=temporal_cutoff\n        )\n        \n        template_predictions[target_id] = preds\n        if n_needed > 0:\n            protenix_queue[target_id] = (n_needed, next_pred, sequence)\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    protenix_predictions = {}\n    \n    if protenix_queue and USE_PROTENIX:\n        print(f\"\\nPHASE 2: Protenix for {len(protenix_queue)} targets\")\n        \n        if protenix_output_path is None:\n            protenix_output_path = Path(OUTPUT_PATH) / f\"{dataset_name}_protenix\"\n        protenix_output_path = Path(protenix_output_path)\n        protenix_output_path.mkdir(parents=True, exist_ok=True)\n        input_path = Path(DATA_PATH)\n        \n        with protenix_context():\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                        target_id, sequence, protenix_output_path, input_path,\n                        seed=101, n_cycle=12, n_sample=n_needed, n_step=250\n                    )\n                    print(f\"    Done in {(time.time()-t0)/60:.1f} min\")\n                    \n                    preds = get_protenix_predictions(\n                        target_id, sequence, protenix_output_path,\n                        seed=101, n_sample=n_needed\n                    )\n                    protenix_predictions[target_id] = preds\n                    print(f\"    Got {len(preds)} Protenix predictions\")\n                    \n                except Exception as e:\n                    print(f\"    Protenix FAILED: {e}\")\n                    protenix_predictions[target_id] = None\n        \n        ptx_time = time.time() - start_time - template_time\n        print(f\"\\nPhase 2 done: {ptx_time/60:.1f} min\")\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    print(f\"\\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        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        n_denovo = 0\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 = enhanced_rna_constraints(de_novo, target_id, segments_map, confidence=0.15, passes=3)\n            predictions.append(refined)\n            pred_num += 1\n            n_denovo += 1\n        \n        if n_denovo > 0:\n            print(f\"  {target_id}: filled {n_denovo} slots with de novo\")\n        \n        for j in range(len(sequence)):\n            pred_row = {\n                'ID': f\"{target_id}_{j+1}\",\n                'resname': sequence[j],\n                'resid': j + 1,\n            }\n            for i in range(5):\n                pred_row[f'x_{i+1}'] = predictions[i][j][0]\n                pred_row[f'y_{i+1}'] = predictions[i][j][1]\n                pred_row[f'z_{i+1}'] = predictions[i][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\"  Ensemble template 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","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"template_info_dict.clear()\nprediction_metadata_dict.clear()\n\ntest_predictions = generate_predictions_batch(\n    test_seqs,\n    combined_seqs,\n    combined_coords_dict,\n    dataset_name=\"test\",\n    use_temporal_cutoff=False,\n)\n\ntest_predictions.to_csv('submission.csv', index=False)\nprint(\"Saved: submission.csv\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}