{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":118765,"databundleVersionId":15231210},{"sourceType":"datasetVersion","sourceId":10880419,"datasetId":6760509,"databundleVersionId":11247150},{"sourceType":"datasetVersion","sourceId":14604295,"datasetId":9328538,"databundleVersionId":15440074},{"sourceType":"datasetVersion","sourceId":10880374,"datasetId":6760482,"databundleVersionId":11247092},{"sourceType":"datasetVersion","sourceId":11899194,"datasetId":7479946,"databundleVersionId":12404228},{"sourceType":"datasetVersion","sourceId":11451236,"datasetId":7174725,"databundleVersionId":11892125},{"sourceType":"datasetVersion","sourceId":11230242,"datasetId":7014687,"databundleVersionId":11640305},{"sourceType":"datasetVersion","sourceId":13282339,"datasetId":7162026,"databundleVersionId":13982667},{"sourceType":"datasetVersion","sourceId":14874339,"datasetId":9502242,"databundleVersionId":15736806},{"sourceType":"datasetVersion","sourceId":10855324,"datasetId":6742586,"databundleVersionId":11219268}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\"\"\"\nV70: Template + Protenix — Fixed V69\n\nBased on V69 with fix: removed ligand context (caused 0 predictions for 8 targets)\n  1. sorted_by_ranking_score=True for Protenix output ordering\n  2. Multi-seed inference (3 seeds x 7 samples = 21 total) for diversity\n  3. Selective n_step=400 for targets >=200nt\n  4. RNA constraints applied to Protenix outputs (not just TBM)\n\nDatasets required:\n  - Stanford RNA 3D Folding 2\n  - protenix-packages (whl + USalign)\n  - protenix_mg_packages\n  - protenix-finetuned-rna3db-all-1599 (checkpoint)\n  - protenix-rmsa-repo (source code)\n  - biopython\n  - ml-collections\n\"\"\"\n\nimport os\nimport sys\n\n# ============================================================\n# Environment setup\n# ============================================================\nIS_KAGGLE = os.path.exists('/kaggle/input/stanford-rna-3d-folding-2/')\n\nif IS_KAGGLE:\n    DATA_PATH = '/kaggle/input/stanford-rna-3d-folding-2/'\n    OUTPUT_PATH = '/kaggle/working/output'\n    USALIGN_BIN = '/kaggle/working/USalign'\n    PROTENIX_DIR = '/kaggle/working/Protenix'\n\n    # Copy USalign binary\n    # Copy USalign binary (try both possible locations)\n    if os.path.exists('/kaggle/input/datasets/metric/usalign'):\n        os.system('cp /kaggle/input/datasets/metric/usalign/USalign /kaggle/working/ 2>/dev/null')\n    elif os.path.exists('/kaggle/input/datasets/zoushuxian/protenix-packages/packages/USalign'):\n        os.system('cp /kaggle/input/datasets/zoushuxian/protenix-packages/packages/USalign /kaggle/working/')\n    os.system('chmod +x /kaggle/working/USalign')\n\n    # Install dependencies from .whl datasets\n    os.system('cp -r /kaggle/input/datasets/zoushuxian/protenix-packages/packages /kaggle/working')\n    os.chdir('/kaggle/working/packages')\n    os.system('pip install --no-deps --exists-action=i *.whl')\n    os.chdir('/kaggle/working')\n\n    # ihm and modelcif from source\n    os.system('mv /kaggle/working/packages/ihm-2.3/ihm-2.3 /kaggle/working 2>/dev/null')\n    os.system('mv /kaggle/working/packages/modelcif-0.7/modelcif-0.7 /kaggle/working 2>/dev/null')\n    os.system('pip install /kaggle/working/ihm-2.3 2>/dev/null')\n    os.system('pip install /kaggle/working/modelcif-0.7 2>/dev/null')\n    os.system('rm -rf /kaggle/working/ihm-2.3 /kaggle/working/modelcif-0.7')\n\n    # Biopython and ml_collections\n    import glob as _gl\n    for whl in _gl.glob('/kaggle/input/datasets/ogurtsov/biopython/*.whl') + _gl.glob('/kaggle/input/datasets/kami1976/biopython-cp312/*.whl'):\n        os.system(f'pip install {whl}')\n    for whl in _gl.glob('/kaggle/input/datasets/ogurtsov/ml-collections/*.whl'):\n        os.system(f'pip install {whl}')\n\n    os.system('rm -rf /kaggle/working/packages')\n\n    # Extra mg packages\n    os.system('cp -r /kaggle/input/datasets/zoushuxian/protenix-mg-packages/protenix_mg_packages /kaggle/working')\n    os.chdir('/kaggle/working/protenix_mg_packages')\n    os.system('pip install --no-deps --exists-action=i *.whl')\n    os.chdir('/kaggle/working')\n    os.system('rm -rf /kaggle/working/protenix_mg_packages')\n\n    # Copy Protenix source\n    os.system('cp -R /kaggle/input/datasets/zoushuxian/protenix-rmsa-repo/protenix_kaggle /kaggle/working/')\n    os.system('mv /kaggle/working/protenix_kaggle /kaggle/working/Protenix')\nelse:\n    raise ValueError(\"This script is designed for Kaggle environment\")\n\nprint(f\"Data path: {DATA_PATH}\")\nprint(f\"Protenix dir: {PROTENIX_DIR}\")\n\nimport json\nimport numpy as np\nimport pandas as pd\nimport time\nimport random\nimport warnings\nimport contextlib\nimport shutil\nfrom pathlib import Path\n\nwarnings.filterwarnings('ignore')\n\ndef seed_everything(seed: int = 42):\n    random.seed(seed)\n    np.random.seed(seed)\n\nseed_everything(42)\n\n# ============================================================\n# Load data\n# ============================================================\nprint(\"Loading sequence data...\")\ntrain_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\nprint(f\"Train: {len(train_seqs)}, Val: {len(validation_seqs)}, Test: {len(test_seqs)}\")\n\n# Combine train+val for test prediction (more templates = better)\ncombined_seqs = pd.concat([train_seqs, validation_seqs], ignore_index=True)\ncombined_labels = pd.concat([train_labels, validation_labels], ignore_index=True)\n\nUSE_PROTENIX = True\nMIN_SIMILARITY = 0.0\nMIN_PERCENT_IDENTITY = 50\n\n# ============================================================\n# Alignment setup\n# ============================================================\nfrom Bio.Align import PairwiseAligner\n\ndef make_aligner():\n    \"\"\"Create a PairwiseAligner with tuned gap penalties for RNA TBM.\"\"\"\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# ============================================================\n# FASTA / stoichiometry / segment helpers\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: 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    if pd.isna(stoich) or pd.isna(all_seq) or str(stoich).strip() == \"\" or str(all_seq).strip() == \"\":\n        return [(0, len(seq))]\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: 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:\n        return [(0, len(seq))]\n\ndef build_segments_map(df):\n    seg_map, stoich_map = {}, {}\n    for _, r in df.iterrows():\n        tid = r['target_id']\n        seg_map[tid] = get_chain_segments(r)\n        stoich_map[tid] = str(r.get('stoichiometry', '') if not pd.isna(r.get('stoichiometry', '')) else '')\n    return seg_map, stoich_map\n\n# ============================================================\n# Data processing\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\ntrain_coords_dict = process_labels(train_labels)\ncombined_coords_dict = process_labels(combined_labels)\n\nprint(f\"Templates: {len(train_coords_dict)} train, {len(combined_coords_dict)} combined\")\n\n# ============================================================\n# Template search\n# ============================================================\ndef find_similar_sequences_detailed(query_seq, train_seqs_df, train_coords_dict,\n                                    temporal_cutoff=None, top_n=5):\n    similar_seqs = []\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        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        percent_identity = 100 * identical / len(query_seq)\n\n        similar_seqs.append((\n            target_id, train_seq, normalized_score,\n            train_coords_dict[target_id], percent_identity\n        ))\n\n    similar_seqs.sort(key=lambda x: x[2], reverse=True)\n    return similar_seqs[:top_n]\n\n\n# ============================================================\n# Coordinate adaptation (sentinel-aware)\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].copy()\n        if len(t_chunk) == (q_end - q_start):\n            sentinel_mask = np.abs(t_chunk[:, 0]) > 1e6\n            if sentinel_mask.any():\n                t_chunk[sentinel_mask] = np.nan\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, 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\n# ============================================================\n# Physics constraints\n# ============================================================\ndef adaptive_rna_constraints(coordinates, target_id, segments_map, confidence=1.0, passes=2):\n    coords = coordinates.copy()\n    segments = segments_map.get(target_id, [(0, len(coords))])\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            if L < 3:\n                coords[s:e] = X\n                continue\n\n            valid = np.abs(X[:, 0]) < 1e6\n\n            d = X[1:] - X[:-1]\n            dist = np.linalg.norm(d, axis=1) + 1e-6\n            scale = (5.95 - dist) / dist\n            adj = (d * scale[:, None]) * (0.22 * strength)\n            bond_valid = valid[:-1] & valid[1:]\n            adj[~bond_valid] = 0\n            X[:-1] -= adj\n            X[1:] += adj\n\n            d2 = X[2:] - X[:-2]\n            dist2 = np.linalg.norm(d2, axis=1) + 1e-6\n            scale2 = (10.2 - dist2) / dist2\n            adj2 = (d2 * scale2[:, None]) * (0.10 * strength)\n            bond2_valid = valid[:-2] & valid[2:]\n            adj2[~bond2_valid] = 0\n            X[:-2] -= adj2\n            X[2:] += adj2\n\n            lap = 0.5 * (X[:-2] + X[2:]) - X[1:-1]\n            lap_valid = valid[:-2] & valid[1:-1] & valid[2:]\n            lap[~lap_valid] = 0\n            X[1:-1] += (0.06 * strength) * lap\n\n            if L >= 25:\n                k = min(L, 160) if L > 220 else L\n                idx = np.linspace(0, L - 1, k).astype(int) if k < L else np.arange(L)\n                idx = idx[valid[idx]]\n                if len(idx) > 5:\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                    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    return coords\n\n\n# ============================================================\n# Diversity transforms\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=25):\n    s, e = seg\n    L = e - s\n    if L < 30: return coords\n    pivot = s + int(rng.integers(10, L - 10))\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 jitter_chains(coords, segments, rng, max_angle_deg=12, max_trans=1.5):\n    X = coords.copy()\n    gc = X.mean(axis=0, keepdims=True)\n    for (s, e) in segments:\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        shift = rng.normal(size=3)\n        shift = shift / (np.linalg.norm(shift) + 1e-12) * float(rng.uniform(0.0, max_trans))\n        c = X[s:e].mean(axis=0, keepdims=True)\n        X[s:e] = (X[s:e] - c) @ R.T + c + shift\n    X -= X.mean(axis=0, keepdims=True) - gc\n    return X\n\ndef smooth_wiggle(coords, segments, rng, amp=0.8):\n    X = coords.copy()\n    for (s, e) in segments:\n        L = e - s\n        if L < 20: continue\n        ctrl_x = np.linspace(0, L - 1, 6)\n        ctrl_disp = rng.normal(0, amp, size=(6, 3))\n        t = np.arange(L)\n        disp = np.vstack([np.interp(t, ctrl_x, ctrl_disp[:, k]) for k in range(3)]).T\n        X[s:e] += disp\n    return X\n\n\n# ============================================================\n# Protenix utilities\n# ============================================================\nfrom biotite.structure.io.pdbx import CIFFile, get_structure\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\n\n# ============================================================\n# Protenix JSON — RNA-only (no ligand context to avoid chain assignment issues)\n# ============================================================\ndef prepare_protenix_json(target_id, sequence, output_path, input_path, max_length=400):\n    \"\"\"Prepare Protenix input JSON with RNA sequence only.\"\"\"\n    if len(sequence) <= max_length:\n        input_json = [{\n            \"sequences\": [{\"rnaSequence\": {\n                \"sequence\": sequence, \"count\": 1,\n                \"msa\": {\"precomputed_msa_dir\": f\"{input_path}/MSA/{target_id}.MSA.fasta\",\n                        \"pairing_db\": \"rnacentral\"}\n            }}],\n            \"name\": target_id,\n        }]\n    else:\n        print(f\"    Sequence too long ({len(sequence)} > {max_length}), no MSA\")\n        input_json = [{\n            \"sequences\": [{\"rnaSequence\": {\"sequence\": sequence[:max_length], \"count\": 1}}],\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\n\n# ============================================================\n# [OPT 1, 2, 4] Protenix inference with multi-seed + ranking + adaptive n_step\n# ============================================================\ndef run_protenix_inference(target_id, sequence, output_path, input_path,\n                           seed=101, n_cycle=10, n_sample=5, n_step=200, max_length=400):\n    checkpoint_path = \"/kaggle/input/datasets/zoushuxian/protenix-finetuned-rna3db-all-1599/1599_ema_0.999.pt\"\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        \"--sorted_by_ranking_score True\",      # [OPT 1] Sort by ranking score\n        \"\",\n    ]\n\n    from runner.inference import run\n    run()\n\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\n# ============================================================\n# Protenix Selection Logic (Best-of-N Centroid)\n# ============================================================\ndef pairwise_rmsd(coords1, coords2):\n    diff = coords1 - coords2\n    return np.sqrt(np.mean(np.sum(diff**2, axis=1)))\n\ndef select_best_protenix(predictions):\n    \"\"\"Select the 'centroid' prediction that is most similar to all others.\"\"\"\n    n = len(predictions)\n    if n < 2:\n        return predictions\n\n    rmsds = np.zeros((n, n))\n    for i in range(n):\n        for j in range(i + 1, n):\n            d = pairwise_rmsd(predictions[i], predictions[j])\n            rmsds[i, j] = d\n            rmsds[j, i] = d\n\n    mean_rmsds = rmsds.mean(axis=1)\n    sorted_indices = np.argsort(mean_rmsds)\n    return [predictions[i] for i in sorted_indices]\n\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\n\n# ============================================================\n# De novo fallback\n# ============================================================\ndef generate_rna_structure(sequence, seed=None):\n    \"\"\"Generate RNA-like backbone with realistic C1'-C1' distances (~5.9A).\"\"\"\n    rng = np.random.default_rng(seed)\n    n = len(sequence)\n    coords = np.zeros((n, 3))\n    for i in range(1, n):\n        if i == 1:\n            direction = np.array([1.0, 0.0, 0.0])\n        else:\n            prev_dir = coords[i-1] - coords[i-2]\n            prev_dir = prev_dir / (np.linalg.norm(prev_dir) + 1e-12)\n            perturb = rng.normal(0, 0.5, 3)\n            direction = prev_dir + perturb\n            direction = direction / (np.linalg.norm(direction) + 1e-12)\n        coords[i] = coords[i-1] + direction * 5.9\n    coords -= coords.mean(axis=0)\n    return coords\n\n\n# ============================================================\n# Sanitize coordinates\n# ============================================================\ndef sanitize_coords(coords):\n    coords = coords.copy()\n    bad = np.abs(coords[:, 0]) > 1e6\n    if not bad.any():\n        return coords\n\n    for i in np.where(bad)[0]:\n        coords[i] = np.nan\n\n    for i in range(len(coords)):\n        if np.isnan(coords[i, 0]):\n            prev_v = next((j for j in range(i - 1, -1, -1) if not np.isnan(coords[j, 0])), -1)\n            next_v = next((j for j in range(i + 1, len(coords)) if not np.isnan(coords[j, 0])), -1)\n            if prev_v >= 0 and next_v >= 0:\n                w = (i - prev_v) / (next_v - prev_v)\n                coords[i] = (1 - w) * coords[prev_v] + w * coords[next_v]\n            elif prev_v >= 0:\n                coords[i] = coords[prev_v] + [3, 0, 0]\n            elif next_v >= 0:\n                coords[i] = coords[next_v] + [3, 0, 0]\n            else:\n                coords[i] = [i * 3, 0, 0]\n\n    return np.nan_to_num(coords)\n\n\n# ============================================================\n# Prediction metadata tracking\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] = {'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\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\n\n# ============================================================\n# Main prediction pipeline\n# ============================================================\ndef predict_with_templates(sequence, target_id, train_seqs_df, train_coords_dict,\n                           segments_map, n_predictions=5, temporal_cutoff=None):\n    \"\"\"Phase 1: sequential template-based predictions.\"\"\"\n    predictions = []\n    pred_num = 1\n\n    print(f\"\\n{'─'*70}\")\n    print(f\"Target: {target_id} ({len(sequence)} nt)\")\n\n    similar_seqs = find_similar_sequences_detailed(\n        sequence, train_seqs_df, train_coords_dict,\n        temporal_cutoff=temporal_cutoff, top_n=n_predictions\n    )\n\n    if similar_seqs:\n        for i, (tmpl_id, tmpl_seq, sim, tmpl_coords, pct_id) in enumerate(similar_seqs):\n\n            if (sim < MIN_SIMILARITY or pct_id < MIN_PERCENT_IDENTITY) and len(tmpl_seq) < 500:\n                print(f\"  Template {i+1}: {tmpl_id} SKIPPED (sim={sim:.3f}, id={pct_id:.1f}%)\")\n                break\n\n            if USE_PROTENIX and len(sequence) < 100 and i == 4:\n                print(f\"  Short sequence, keeping 1 slot for Protenix\")\n                break\n\n            record_template_info(target_id, tmpl_id, sim, pct_id)\n            record_prediction_metadata(target_id, pred_num, 'template', tmpl_id, sim, pct_id)\n            print(f\"  Template {i+1}: {tmpl_id} (sim={sim:.3f}, id={pct_id:.1f}%)\")\n\n            adapted = adapt_template_to_query(sequence, tmpl_seq, tmpl_coords)\n            refined = adaptive_rna_constraints(adapted, target_id, segments_map,\n                                               confidence=sim)\n            predictions.append(refined)\n            pred_num += 1\n\n            if len(predictions) >= n_predictions:\n                break\n\n    n_needed = n_predictions - len(predictions)\n    if n_needed > 0:\n        print(f\"  → {len(predictions)} from templates, {n_needed} slots for Protenix\")\n    else:\n        print(f\"  → All {n_predictions} from templates\")\n\n    return predictions, n_needed, pred_num\n\n\n\n\n\n# ============================================================\n# Batch prediction pipeline\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    # ---- 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        tc = row.get('temporal_cutoff', None) if use_temporal_cutoff else None\n\n        preds, n_needed, next_pred = predict_with_templates(\n            sequence, target_id, train_seqs_df, train_coords_dict,\n            segments_map, n_predictions=5, temporal_cutoff=tc\n        )\n\n        template_predictions[target_id] = preds\n        if n_needed > 0:\n            protenix_queue[target_id] = (n_needed, next_pred, sequence, row)\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        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, row)) in enumerate(protenix_queue.items()):\n                print(f\"\\n  [{i+1}/{len(protenix_queue)}] {target_id} \"\n                      f\"({len(sequence)} nt, need {n_needed})\")\n\n                try:\n                    prepare_protenix_json(\n                        target_id, sequence, protenix_output_path, input_path\n                    )\n\n                    t0 = time.time()\n\n                    # [OPT 4] Adaptive n_step: 400 for longer targets\n                    n_step = 400 if len(sequence) >= 200 else 200\n\n                    # [OPT 2] Multi-seed inference: 3 seeds x 7 samples = 21 total\n                    all_preds = []\n                    for ptx_seed in [101, 201, 301]:\n                        run_protenix_inference(\n                            target_id, sequence, protenix_output_path, input_path,\n                            seed=ptx_seed, n_cycle=10, n_sample=7, n_step=n_step\n                        )\n                        seed_preds = get_protenix_predictions(\n                            target_id, sequence, protenix_output_path,\n                            seed=ptx_seed, n_sample=7\n                        )\n                        all_preds.extend(seed_preds)\n\n                    print(f\"    Done in {(time.time()-t0)/60:.1f} min, got {len(all_preds)} predictions\")\n\n                    # Select best by centroid RMSD\n                    if len(all_preds) > 1:\n                        all_preds = select_best_protenix(all_preds)\n                        print(f\"    Selected best centroid from {len(all_preds)} samples\")\n\n                    protenix_predictions[target_id] = all_preds\n\n                except Exception as e:\n                    print(f\"    Protenix FAILED: {e}\")\n                    import traceback\n                    traceback.print_exc()\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\")\n\n    # ---- Phase 3: Combine and format ----\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        # 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                    # [OPT 5] Apply RNA constraints to Protenix outputs\n                    refined = adaptive_rna_constraints(\n                        coords, target_id, segments_map, confidence=0.85\n                    )\n                    predictions.append(refined)\n                    pred_num += 1\n                    if len(predictions) >= 5:\n                        break\n\n        # Fill remaining with de novo\n        n_denovo = 0\n        while len(predictions) < 5:\n            record_prediction_metadata(target_id, pred_num, 'de_novo')\n            seed_val = row.name * 1000000 + 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_denovo += 1\n\n        if n_denovo > 0:\n            print(f\"  {target_id}: filled {n_denovo} slots with de novo\")\n\n        # Sanitize predictions\n        for k in range(len(predictions)):\n            predictions[k] = sanitize_coords(predictions[k])\n\n        # Format output rows\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 k in range(5):\n                pred_row[f'x_{k+1}'] = predictions[k][j][0]\n                pred_row[f'y_{k+1}'] = predictions[k][j][1]\n                pred_row[f'z_{k+1}'] = predictions[k][j][2]\n            all_rows.append(pred_row)\n\n    submission_df = pd.DataFrame(all_rows)\n    col_order = ['ID', 'resname', 'resid']\n    for k in range(1, 6):\n        for c in ['x', 'y', 'z']:\n            col_order.append(f'{c}_{k}')\n    submission_df = submission_df[col_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\n\n\n# ============================================================\n# RUN: Test submission\n# ============================================================\ntemplate_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\")\nprint(test_predictions.head())\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T06:12:38.124624Z","iopub.execute_input":"2026-02-24T06:12:38.124975Z","iopub.status.idle":"2026-02-24T06:12:41.018344Z","shell.execute_reply.started":"2026-02-24T06:12:38.124941Z","shell.execute_reply":"2026-02-24T06:12:41.016965Z"}},"outputs":[],"execution_count":null}]}