{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.0"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":118765,"databundleVersionId":15231210},{"sourceType":"datasetVersion","sourceId":14604295,"datasetId":9328538,"databundleVersionId":15440074},{"sourceType":"datasetVersion","sourceId":14874339,"datasetId":9502242,"databundleVersionId":15736806},{"sourceType":"datasetVersion","sourceId":10855324,"datasetId":6742586,"databundleVersionId":11219268}],"dockerImageVersionId":31260,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Stanford RNA 3D Folding Part 2 — Recovery Pipeline\n\nBased on: stable-0-413-without-hash (3).ipynb (LB ~0.408)\n\n## VERSION History\n| version | description | LB |\n|---------|-------------|----|\n| 3 | protenix+TBM (original) | 0.408 |\n| 8 | wrong constraints + ensemble | 0.342 ❌ |\n| 9 | over-constraining | 0.369 ❌ |\n| 10 | BASELINE RESTORE | target: 0.408 |\n| 11 | + safe improvements (RNA MSA, end-same, jitter) | target: 0.408+ |\n| 12 | + nucleotide-aware constraints | target: above v11 |\n\nChange `VERSION` variable in the code cell to select version.","metadata":{}},{"cell_type":"code","source":"!pip install /kaggle/input/datasets/kami1976/biopython-cp312/biopython-1.86-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install biotite\n!pip install rdkit\n!pip install biopython","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport sys\nimport pandas as pd\n\nIS_KAGGLE = bool(os.environ.get(\"KAGGLE_IS_COMPETITION_RERUN\", \"\"))\nLOCAL_N_SAMPLES = 2\n\nif IS_KAGGLE:\n    print(\"Running in KAGGLE COMPETITION mode — all test targets will be processed.\")\nelse:\n    print(f\"Running in LOCAL mode — only the first {LOCAL_N_SAMPLES} test targets \"\n          f\"will be processed to save time.\")\n\n# %%  ─── Cell 3: Imports & Config ───","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import gc\nimport json\nimport time\n\nos.environ[\"LAYERNORM_TYPE\"] = \"torch\"\nos.environ.setdefault(\"RNA_MSA_DEPTH_LIMIT\", \"512\")\n\nfrom pathlib import Path\n\nimport numpy as np\nimport torch\nfrom Bio.Align import PairwiseAligner\nfrom tqdm import tqdm\n\n\n# ═══════════════════════════════════════════════════════════════════\n# VERSION SELECTOR\n# ═══════════════════════════════════════════════════════════════════\n# 10 = Baseline restore (faithful to original 0.408 notebook)\n# 11 = Safe improvements only (RNA MSA, end-same padding, jitter)\n# 12 = + Nucleotide-aware constraint tuning\nVERSION = 10\nprint(f\"Pipeline VERSION: {VERSION}\")\n\n# ─────────────── Paths & Constants ───────────────────────────────────────────\nDATA_BASE              = \"/kaggle/input/stanford-rna-3d-folding-2\"\nDEFAULT_TEST_CSV       = f\"{DATA_BASE}/test_sequences.csv\"\nDEFAULT_TRAIN_CSV      = f\"{DATA_BASE}/train_sequences.csv\"\nDEFAULT_TRAIN_LBLS     = f\"{DATA_BASE}/train_labels.csv\"\nDEFAULT_VAL_CSV        = f\"{DATA_BASE}/validation_sequences.csv\"\nDEFAULT_VAL_LBLS       = f\"{DATA_BASE}/validation_labels.csv\"\nDEFAULT_OUTPUT         = \"/kaggle/working/submission.csv\"\n\nDEFAULT_CODE_DIR = (\n    \"/kaggle/input/datasets/qiweiyin/protenix-v1-adjusted\"\n    \"/Protenix-v1-adjust-v2/Protenix-v1-adjust-v2/Protenix-v1\"\n)\nDEFAULT_ROOT_DIR = DEFAULT_CODE_DIR\n\nMODEL_NAME    = \"protenix_base_20250630_v1.0.0\"\nN_SAMPLE      = 5\nSEED          = 42\nMAX_SEQ_LEN   = int(os.environ.get(\"MAX_SEQ_LEN\",   \"512\"))\nCHUNK_OVERLAP = int(os.environ.get(\"CHUNK_OVERLAP\",  \"64\"))\n\n# ─── Constraint parameters (VERSION-dependent) ───\n# ORIGINAL VALUES that produced 0.408:\nBOND_TARGET_I_I1   = 5.95\nBOND_TARGET_I_I2   = 10.2\nSELF_AVOID_DIST    = 3.2\nCONSTRAINT_PASSES  = 2\nCONSTRAINT_BOND_STRENGTH = 0.22\n\n# V12: Per-dinucleotide bond targets (conservative: close to 5.95)\nDINUC_BOND_TARGETS = {\n    ('C', 'G'): 5.70,  # EDA median 5.53 → conservative\n    ('C', 'U'): 5.70,  # EDA median 5.56 → conservative\n    ('G', 'A'): 6.10,  # EDA median 6.02\n    ('U', 'A'): 6.40,  # EDA median 6.76 → conservative\n    ('U', 'U'): 6.10,  # EDA median 6.02\n    ('G', 'G'): 5.80,  # EDA median 5.60\n}\nDINUC_DEFAULT = 5.95  # Fallback: original value\n\n# TBM quality thresholds\nMIN_SIMILARITY       = float(os.environ.get(\"MIN_SIMILARITY\",       \"0.30\"))\nMIN_PERCENT_IDENTITY = float(os.environ.get(\"MIN_PERCENT_IDENTITY\", \"50.0\"))\n\nUSE_PROTENIX = True\n\n\ndef parse_bool(value, default=False):\n    v = str(value).strip().lower()\n    if v in {\"1\", \"true\", \"t\", \"yes\", \"y\", \"on\"}:\n        return \"true\"\n    if v in {\"0\", \"false\", \"f\", \"no\", \"n\", \"off\"}:\n        return \"false\"\n    return \"true\" if default else \"false\"\n\n\nUSE_MSA      = parse_bool(os.environ.get(\"USE_MSA\",      \"false\"))\nUSE_TEMPLATE = parse_bool(os.environ.get(\"USE_TEMPLATE\", \"false\"))\n\n# V11+: Enable RNA MSA (shown effective in Part 1 solutions)\nif VERSION >= 11:\n    USE_RNA_MSA = parse_bool(os.environ.get(\"USE_RNA_MSA\", \"true\"), default=True)\nelse:\n    USE_RNA_MSA = parse_bool(os.environ.get(\"USE_RNA_MSA\", \"false\"))\n\nMODEL_N_SAMPLE = int(os.environ.get(\"MODEL_N_SAMPLE\", str(N_SAMPLE)))\n\n\n# ─────────────── General Utilities ───────────────────────────────────────────\ndef seed_everything(seed):\n    os.environ[\"CUBLAS_WORKSPACE_CONFIG\"] = \":4096:8\"\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    np.random.seed(seed)\n    torch.backends.cudnn.benchmark = False\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.enabled = True\n    torch.use_deterministic_algorithms(True)\n\n\ndef resolve_paths():\n    test_csv   = os.environ.get(\"TEST_CSV\",           DEFAULT_TEST_CSV)\n    output_csv = os.environ.get(\"SUBMISSION_CSV\",     DEFAULT_OUTPUT)\n    code_dir   = os.environ.get(\"PROTENIX_CODE_DIR\",  DEFAULT_CODE_DIR)\n    root_dir   = os.environ.get(\"PROTENIX_ROOT_DIR\",  DEFAULT_ROOT_DIR)\n    return test_csv, output_csv, code_dir, root_dir\n\n\ndef ensure_required_files(root_dir):\n    for p, name in [\n        (Path(root_dir) / \"checkpoint\" / f\"{MODEL_NAME}.pt\",          \"checkpoint\"),\n        (Path(root_dir) / \"common\" / \"components.cif\",                \"CCD file\"),\n        (Path(root_dir) / \"common\" / \"components.cif.rdkit_mol.pkl\",  \"CCD cache\"),\n    ]:\n        if not p.exists():\n            raise FileNotFoundError(f\"Missing {name}: {p}\")\n\n\n# ─────────────── Protenix Input / Config Helpers ─────────────────────────────\ndef build_input_json(df, json_path):\n    data = [\n        {\n            \"name\": row[\"target_id\"],\n            \"covalent_bonds\": [],\n            \"sequences\": [{\"rnaSequence\": {\"sequence\": row[\"sequence\"], \"count\": 1}}],\n        }\n        for _, row in df.iterrows()\n    ]\n    with open(json_path, \"w\", encoding=\"utf-8\") as f:\n        json.dump(data, f)\n\n\ndef build_configs(input_json_path, dump_dir, model_name):\n    from configs.configs_base import configs as configs_base\n    from configs.configs_data import data_configs\n    from configs.configs_inference import inference_configs\n    from configs.configs_model_type import model_configs\n    from protenix.config.config import parse_configs\n\n    base = {**configs_base, **{\"data\": data_configs}, **inference_configs}\n\n    def deep_update(t, p):\n        for k, v in p.items():\n            if isinstance(v, dict) and k in t and isinstance(t[k], dict):\n                deep_update(t[k], v)\n            else:\n                t[k] = v\n\n    deep_update(base, model_configs[model_name])\n    arg_str = \" \".join([\n        f\"--model_name {model_name}\",\n        f\"--input_json_path {input_json_path}\",\n        f\"--dump_dir {dump_dir}\",\n        f\"--use_msa {USE_MSA}\",\n        f\"--use_template {USE_TEMPLATE}\",\n        f\"--use_rna_msa {USE_RNA_MSA}\",\n        f\"--sample_diffusion.N_sample {MODEL_N_SAMPLE}\",\n        f\"--seeds {SEED}\",\n    ])\n    return parse_configs(configs=base, arg_str=arg_str, fill_required_with_null=True)\n\n\ndef get_c1_mask(data, atom_array):\n    if atom_array is not None:\n        try:\n            if hasattr(atom_array, \"centre_atom_mask\"):\n                m = atom_array.centre_atom_mask == 1\n                if hasattr(atom_array, \"is_rna\"):\n                    m = m & atom_array.is_rna\n                return torch.from_numpy(m).bool()\n            if hasattr(atom_array, \"atom_name\"):\n                base = atom_array.atom_name == \"C1'\"\n                if hasattr(atom_array, \"is_rna\"):\n                    base = base & atom_array.is_rna\n                return torch.from_numpy(base).bool()\n        except Exception:\n            pass\n    f = data[\"input_feature_dict\"]\n    if \"center_atom_mask\" in f:\n        return (f[\"center_atom_mask\"] == 1).bool()\n    if \"centre_atom_mask\" in f:\n        return (f[\"centre_atom_mask\"] == 1).bool()\n    return (f[\"atom_to_tokatom_idx\"] == 11).bool()\n\n\ndef get_feature_c1_mask(data):\n    f = data[\"input_feature_dict\"]\n    if \"centre_atom_mask\" in f:\n        return f[\"centre_atom_mask\"].long() == 1\n    return f[\"atom_to_tokatom_idx\"].long() == 12\n\n\ndef coords_to_rows(target_id, seq, coords):\n    rows = []\n    for i in range(len(seq)):\n        row = {\"ID\": f\"{target_id}_{i + 1}\", \"resname\": seq[i], \"resid\": i + 1}\n        for s in range(N_SAMPLE):\n            if s < coords.shape[0] and i < coords.shape[1]:\n                x, y, z = coords[s, i]\n            else:\n                x, y, z = 0.0, 0.0, 0.0\n            row[f\"x_{s + 1}\"] = float(x)\n            row[f\"y_{s + 1}\"] = float(y)\n            row[f\"z_{s + 1}\"] = float(z)\n        rows.append(row)\n    return rows\n\n\ndef pad_samples(coords, n):\n    if coords.shape[0] >= n:\n        return coords[:n]\n    if coords.shape[0] == 0:\n        return np.zeros((n, coords.shape[1], 3), dtype=coords.dtype)\n    extra = np.repeat(coords[:1], n - coords.shape[0], axis=0)\n    return np.concatenate([coords, extra], axis=0)\n\n\n# ─────────────── TBM Core Functions ──────────────────────────────────────────\ndef _make_aligner():\n    al = PairwiseAligner()\n    al.mode                           = \"global\"\n    al.match_score                    = 2\n    al.mismatch_score                 = -1.5\n    al.open_gap_score                 = -8\n    al.extend_gap_score               = -0.4\n    al.query_left_open_gap_score      = -8\n    al.query_left_extend_gap_score    = -0.4\n    al.query_right_open_gap_score     = -8\n    al.query_right_extend_gap_score   = -0.4\n    al.target_left_open_gap_score     = -8\n    al.target_left_extend_gap_score   = -0.4\n    al.target_right_open_gap_score    = -8\n    al.target_right_extend_gap_score  = -0.4\n    return al\n\n_aligner = _make_aligner()\n\n\ndef parse_stoichiometry(stoich):\n    if pd.isna(stoich) or str(stoich).strip() == \"\":\n        return []\n    return [(ch.strip(), int(cnt)) for part in str(stoich).split(\";\")\n            for ch, cnt in [part.split(\":\")]]\n\n\ndef parse_fasta(fasta_content):\n    out, cur, parts = {}, None, []\n    for line in str(fasta_content).splitlines():\n        line = line.strip()\n        if not line:\n            continue\n        if line.startswith(\">\"):\n            if cur is not None:\n                out[cur] = \"\".join(parts)\n            cur = line[1:].split()[0]\n            parts = []\n        else:\n            parts.append(line.replace(\" \", \"\"))\n    if cur is not None:\n        out[cur] = \"\".join(parts)\n    return out\n\n\ndef get_chain_segments(row):\n    seq    = row[\"sequence\"]\n    stoich = row.get(\"stoichiometry\", \"\")\n    all_sq = row.get(\"all_sequences\", \"\")\n    if (pd.isna(stoich) or pd.isna(all_sq)\n            or str(stoich).strip() == \"\" or str(all_sq).strip() == \"\"):\n        return [(0, len(seq))]\n    try:\n        chain_dict = parse_fasta(all_sq)\n        order = parse_stoichiometry(stoich)\n        segs, pos = [], 0\n        for ch, cnt in order:\n            base = chain_dict.get(ch)\n            if base is None:\n                return [(0, len(seq))]\n            for _ in range(cnt):\n                segs.append((pos, pos + len(base)))\n                pos += len(base)\n        return segs if pos == len(seq) else [(0, len(seq))]\n    except Exception:\n        return [(0, len(seq))]\n\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        raw_s             = r.get(\"stoichiometry\", \"\")\n        stoich_map[tid]   = \"\" if pd.isna(raw_s) else str(raw_s)\n    return seg_map, stoich_map\n\n\ndef process_labels(labels_df):\n    coords = {}\n    prefixes = labels_df[\"ID\"].str.rsplit(\"_\", n=1).str[0]\n    for prefix, grp in labels_df.groupby(prefixes):\n        coords[prefix] = grp.sort_values(\"resid\")[[\"x_1\", \"y_1\", \"z_1\"]].values\n    return coords\n\n\ndef _build_aligned_strings(query_seq, template_seq, alignment):\n    q_segs, t_segs = alignment.aligned\n    aq, at, qi, ti = [], [], 0, 0\n    for (qs, qe), (ts, te) in zip(q_segs, t_segs):\n        while qi < qs: aq.append(query_seq[qi]);    at.append(\"-\");              qi += 1\n        while ti < ts: aq.append(\"-\");              at.append(template_seq[ti]); ti += 1\n        for qp, tp in zip(range(qs, qe), range(ts, te)):\n            aq.append(query_seq[qp]); at.append(template_seq[tp])\n        qi, ti = qe, te\n    while qi < len(query_seq):    aq.append(query_seq[qi]);    at.append(\"-\");              qi += 1\n    while ti < len(template_seq): aq.append(\"-\");              at.append(template_seq[ti]); ti += 1\n    return \"\".join(aq), \"\".join(at)\n\n\ndef find_similar_sequences_detailed(query_seq, train_seqs_df, train_coords_dict, top_n=30):\n    \"\"\"Find similar sequences — IDENTICAL to original 0.408 notebook.\"\"\"\n    results = []\n    query_len = len(query_seq)\n\n    if 'seq_len' not in train_seqs_df.columns:\n        train_seqs_df['seq_len'] = train_seqs_df['sequence'].str.len()\n\n    # Original length filter: ±30%\n    min_len = query_len * 0.7\n    max_len = query_len * 1.3\n\n    candidates = train_seqs_df[\n        (train_seqs_df['target_id'].isin(train_coords_dict.keys())) &\n        (train_seqs_df['seq_len'] >= min_len) &\n        (train_seqs_df['seq_len'] <= max_len)\n    ]\n\n    for _, row in candidates.iterrows():\n        tid, tseq = row[\"target_id\"], row[\"sequence\"]\n        aln = next(iter(_aligner.align(query_seq, tseq)))\n        norm_s = aln.score / (2 * min(query_len, len(tseq)))\n\n        if norm_s < 0.2:\n            continue\n\n        identical = sum(\n            1 for (qs, qe), (ts, te) in zip(*aln.aligned)\n            for qp, tp in zip(range(qs, qe), range(ts, te))\n            if query_seq[qp] == tseq[tp]\n        )\n        pct_id = 100 * identical / query_len\n\n        aq, at = _build_aligned_strings(query_seq, tseq, aln)\n        results.append((tid, tseq, norm_s, train_coords_dict[tid], pct_id, aq, at))\n\n    results.sort(key=lambda x: x[2], reverse=True)\n    return results[:top_n]\n\n\ndef adapt_template_to_query(query_seq, template_seq, template_coords):\n    aln        = next(iter(_aligner.align(query_seq, template_seq)))\n    new_coords = np.full((len(query_seq), 3), np.nan)\n    for (qs, qe), (ts, te) in zip(*aln.aligned):\n        chunk = template_coords[ts:te]\n        if len(chunk) == (qe - qs):\n            new_coords[qs:qe] = chunk\n    for i in range(len(new_coords)):\n        if np.isnan(new_coords[i, 0]):\n            pv = next((j for j in range(i - 1, -1, -1) if not np.isnan(new_coords[j, 0])), -1)\n            nv = next((j for j in range(i + 1, len(new_coords)) if not np.isnan(new_coords[j, 0])), -1)\n            if pv >= 0 and nv >= 0:\n                w = (i - pv) / (nv - pv)\n                new_coords[i] = (1 - w) * new_coords[pv] + w * new_coords[nv]\n            elif pv >= 0:\n                new_coords[i] = new_coords[pv] + [3, 0, 0]\n            elif nv >= 0:\n                new_coords[i] = new_coords[nv] + [3, 0, 0]\n            else:\n                new_coords[i] = [i * 3, 0, 0]\n    return np.nan_to_num(new_coords)\n\n\n# ─────────────── Geometry Constraints ────────────────────────────────────────\ndef adaptive_rna_constraints(coords, target_id, segments_map, confidence=1.0,\n                              passes=None, seq=None):\n    \"\"\"\n    Geometry constraints applied to predicted coordinates.\n    V10: Identical to original 0.408 notebook (5.95Å, 2 passes, 0.22 strength)\n    V12: Optional nucleotide-aware bond targets\n    \"\"\"\n    if passes is None:\n        passes = CONSTRAINT_PASSES\n\n    X        = coords.copy()\n    segments = segments_map.get(target_id, [(0, len(X))])\n    strength = max(0.75 * (1.0 - min(confidence, 0.97)), 0.02)\n\n    for _ in range(passes):\n        for s, e in segments:\n            C = X[s:e]; L = e - s\n            if L < 3:\n                continue\n\n            # bond i–i+1\n            d    = C[1:] - C[:-1]\n            dist = np.linalg.norm(d, axis=1) + 1e-6\n\n            if VERSION >= 12 and seq is not None:\n                # V12: Per-dinucleotide bond targets\n                seg_seq = seq[s:e]\n                targets = np.array([\n                    DINUC_BOND_TARGETS.get((seg_seq[j], seg_seq[j+1]), DINUC_DEFAULT)\n                    for j in range(L - 1)\n                ])\n            else:\n                # V10/V11: Original uniform target\n                targets = BOND_TARGET_I_I1\n\n            adj  = d * ((targets - dist) / dist)[:, None] * (CONSTRAINT_BOND_STRENGTH * strength)\n            C[:-1] -= adj; C[1:] += adj\n\n            # soft i–i+2\n            d2   = C[2:] - C[:-2]; d2n = np.linalg.norm(d2, axis=1) + 1e-6\n            adj2 = d2 * ((BOND_TARGET_I_I2 - d2n) / d2n)[:, None] * (0.10 * strength)\n            C[:-2] -= adj2; C[2:] += adj2\n\n            # Laplacian smoothing\n            C[1:-1] += (0.06 * strength) * (0.5 * (C[:-2] + C[2:]) - C[1:-1])\n\n            # self-avoidance\n            if L >= 25:\n                idx  = np.linspace(0, L - 1, min(L, 160)).astype(int) if L > 220 else np.arange(L)\n                P    = C[idx]; diff = P[:, None, :] - P[None, :, :]\n                dm   = np.linalg.norm(diff, axis=2) + 1e-6\n                sep  = np.abs(idx[:, None] - idx[None, :])\n                mask = (sep > 2) & (dm < SELF_AVOID_DIST)\n                if np.any(mask):\n                    vec = (diff * ((SELF_AVOID_DIST - dm) / dm)[:, :, None] * mask[:, :, None]).sum(axis=1)\n                    C[idx] += (0.015 * strength) * vec\n            X[s:e] = C\n    return X\n\n\ndef _rotmat(axis, ang):\n    a = np.asarray(axis, float); a /= np.linalg.norm(a) + 1e-12\n    x, y, z = a; c, s = np.cos(ang), np.sin(ang); CC = 1 - c\n    return np.array([[c+x*x*CC, x*y*CC-z*s, x*z*CC+y*s],\n                     [y*x*CC+z*s, c+y*y*CC, y*z*CC-x*s],\n                     [z*x*CC-y*s, z*y*CC+x*s, c+z*z*CC]])\n\n\ndef apply_hinge(coords, seg, rng, deg=22):\n    s, e = seg; L = e - s\n    if L < 30: return coords\n    pivot = s + int(rng.integers(10, L - 10))\n    R = _rotmat(rng.normal(size=3), np.deg2rad(float(rng.uniform(-deg, deg))))\n    X = coords.copy(); p0 = X[pivot].copy()\n    X[pivot+1:e] = (X[pivot+1:e] - p0) @ R.T + p0\n    return X\n\n\n# V11+: Increased jitter translation range (Deep EDA: inter-chain median = 10.9Å)\nJITTER_TRANS = 3.0 if VERSION >= 11 else 1.5\n\ndef jitter_chains(coords, segs, rng, deg=12, trans=None):\n    if trans is None:\n        trans = JITTER_TRANS\n    X = coords.copy(); gc_ = X.mean(0, keepdims=True)\n    for s, e in segs:\n        R     = _rotmat(rng.normal(size=3), np.deg2rad(float(rng.uniform(-deg, deg))))\n        shift = rng.normal(size=3); shift = shift / (np.linalg.norm(shift) + 1e-12) * float(rng.uniform(0, trans))\n        c     = X[s:e].mean(0, keepdims=True)\n        X[s:e] = (X[s:e] - c) @ R.T + c + shift\n    X -= X.mean(0, keepdims=True) - gc_\n    return X\n\n\ndef smooth_wiggle(coords, segs, rng, amp=0.8):\n    X = coords.copy()\n    for s, e in segs:\n        L = e - s\n        if L < 20: continue\n        ctrl = np.linspace(0, L - 1, 6); disp = rng.normal(0, amp, (6, 3)); t = np.arange(L)\n        X[s:e] += np.vstack([np.interp(t, ctrl, disp[:, k]) for k in range(3)]).T\n    return X\n\n\ndef generate_rna_structure(sequence, seed=None):\n    \"\"\"De-novo fallback — identical to original for V10, improved for V11+.\"\"\"\n    if seed is not None:\n        np.random.seed(seed)\n    n = len(sequence)\n    coords = np.zeros((n, 3))\n\n    if VERSION >= 11:\n        # V11+: Base-type variation\n        rise = 2.8; radius = 10.0; twist = 0.57\n        offsets = {'A': 0.1, 'C': -0.1, 'G': 0.15, 'U': -0.15}\n        for i in range(n):\n            ang = i * twist\n            off = offsets.get(sequence[i], 0.0)\n            coords[i] = [(radius + off) * np.cos(ang),\n                         (radius + off) * np.sin(ang),\n                         i * rise]\n    else:\n        # V10: Original A-form helix\n        for i in range(n):\n            ang = i * 0.6\n            coords[i] = [10.0 * np.cos(ang), 10.0 * np.sin(ang), i * 2.5]\n    return coords\n\n\n# ─────────────── TBM Phase ───────────────────────────────────────────────────\ndef tbm_phase(test_df, train_seqs_df, train_coords_dict, segments_map):\n    \"\"\"\n    Phase 1 — Template-Based Modeling.\n    V10: IDENTICAL to original 0.408 notebook.\n    V11+: Same TBM logic, improvements only in constraint/diversity params.\n    \"\"\"\n    print(f\"\\n{'='*60}\")\n    print(f\"PHASE 1: Template-Based Modeling (VERSION={VERSION})\")\n    print(f\"  MIN_SIMILARITY = {MIN_SIMILARITY}  |  MIN_PCT_IDENTITY = {MIN_PERCENT_IDENTITY}\")\n    print(f\"  Bond constraint: {BOND_TARGET_I_I1}Å | Passes: {CONSTRAINT_PASSES}\")\n    print(f\"{'='*60}\")\n    t0 = time.time()\n\n    template_predictions = {}\n    protenix_queue       = {}\n\n    for _, row in tqdm(test_df.iterrows(), total=len(test_df), desc=\"TBM Phase\"):\n        tid = row[\"target_id\"]\n        seq = row[\"sequence\"]\n        segs = segments_map.get(tid, [(0, len(seq))])\n\n        similar = find_similar_sequences_detailed(seq, train_seqs_df, train_coords_dict, top_n=30)\n        preds   = []\n        used    = set()\n\n        for i, (tmpl_id, tmpl_seq, sim, tmpl_coords, pct_id, _, _) in enumerate(similar):\n            if len(preds) >= N_SAMPLE:\n                break\n            if sim < MIN_SIMILARITY or pct_id < MIN_PERCENT_IDENTITY:\n                break\n            if tmpl_id in used:\n                continue\n\n            rng     = np.random.default_rng((row.name * 10000000000 + i * 10007) % (2**32))\n            adapted = adapt_template_to_query(seq, tmpl_seq, tmpl_coords)\n\n            # Diversity transforms — IDENTICAL to original 0.408 notebook\n            slot = len(preds)\n            if slot == 0:\n                X = adapted\n            elif slot == 1:\n                X = adapted + rng.normal(0, max(0.01, (0.40 - sim) * 0.06), adapted.shape)\n            elif slot == 2:\n                longest = max(segs, key=lambda se: se[1] - se[0])\n                X = apply_hinge(adapted, longest, rng)\n            elif slot == 3:\n                X = jitter_chains(adapted, segs, rng)\n            else:\n                X = smooth_wiggle(adapted, segs, rng)\n\n            # V12 passes seq for nucleotide-aware constraints\n            refined = adaptive_rna_constraints(X, tid, segments_map, confidence=sim,\n                                                seq=seq if VERSION >= 12 else None)\n            preds.append(refined)\n            used.add(tmpl_id)\n\n        template_predictions[tid] = preds\n        n_needed = N_SAMPLE - len(preds)\n        if n_needed > 0:\n            protenix_queue[tid] = (n_needed, seq)\n            print(f\"  {tid} ({len(seq)} nt): {len(preds)} TBM → need {n_needed} from Protenix\")\n        else:\n            print(f\"  {tid} ({len(seq)} nt): all {N_SAMPLE} from TBM ✓\")\n\n    elapsed = time.time() - t0\n    n_full  = len(test_df) - len(protenix_queue)\n    print(f\"\\nPhase 1 done in {elapsed:.1f}s\")\n    print(f\"  Fully covered by TBM : {n_full}\")\n    print(f\"  Need Protenix        : {len(protenix_queue)}\")\n    return template_predictions, protenix_queue\n\n\n# ─────────────── Main Execution Flow ─────────────────────────────────────────\ndef main():\n    test_csv, output_csv, code_dir, root_dir = resolve_paths()\n\n    if not os.path.isdir(code_dir):\n        raise FileNotFoundError(\n            f\"Missing PROTENIX_CODE_DIR: {code_dir}. \"\n            \"Set PROTENIX_CODE_DIR to the repo path.\"\n        )\n\n    os.environ[\"PROTENIX_ROOT_DIR\"] = root_dir\n    sys.path.append(code_dir)\n    ensure_required_files(root_dir)\n    seed_everything(SEED)\n\n    IS_KAGGLE_RUN = bool(os.environ.get(\"KAGGLE_IS_COMPETITION_RERUN\", \"\"))\n    test_df_full = pd.read_csv(test_csv)\n    test_df = (test_df_full.head(LOCAL_N_SAMPLES) if not IS_KAGGLE_RUN else test_df_full).reset_index(drop=True)\n    print(f\"Test targets : {len(test_df)}\")\n\n    test_df_trunc = test_df.copy()\n    test_df_trunc[\"sequence\"] = test_df_trunc[\"sequence\"].str[:MAX_SEQ_LEN]\n\n    # ── Load TBM data ──\n    print(\"\\nLoading training data for TBM …\")\n    train_seqs = pd.read_csv(DEFAULT_TRAIN_CSV)\n    val_seqs = pd.read_csv(DEFAULT_VAL_CSV)\n    train_labels = pd.read_csv(DEFAULT_TRAIN_LBLS)\n    val_labels = pd.read_csv(DEFAULT_VAL_LBLS)\n\n    combined_seqs = pd.concat([train_seqs, val_seqs], ignore_index=True)\n    combined_labels = pd.concat([train_labels, val_labels], ignore_index=True)\n\n    train_coords = process_labels(combined_labels)\n    segments_map, _ = build_segments_map(test_df)\n\n    print(f\"Template pool: {len(combined_seqs)} sequences, {len(train_coords)} structures\")\n\n    # ─── PHASE 1: TBM ───\n    template_preds, protenix_queue = tbm_phase(\n        test_df, combined_seqs, train_coords, segments_map\n    )\n\n    # ─── PHASE 2: Protenix ───\n    protenix_preds = {}\n\n    if protenix_queue and USE_PROTENIX:\n        print(f\"\\n{'='*60}\")\n        print(f\"PHASE 2: Protenix for {len(protenix_queue)} targets\")\n        print(f\"{'='*60}\")\n\n        work_dir = Path(\"/kaggle/working\")\n        work_dir.mkdir(parents=True, exist_ok=True)\n\n        queue_df = (test_df_trunc[test_df_trunc[\"target_id\"].isin(protenix_queue)]\n                    .reset_index(drop=True))\n        input_json_path = str(work_dir / \"protenix_queue_input.json\")\n        build_input_json(queue_df, input_json_path)\n\n        try:\n            from protenix.data.inference.infer_dataloader import InferenceDataset\n            from runner.inference import (InferenceRunner,\n                                          update_gpu_compatible_configs,\n                                          update_inference_configs)\n\n            configs = build_configs(input_json_path, str(work_dir / \"outputs\"), MODEL_NAME)\n            configs = update_gpu_compatible_configs(configs)\n            runner = InferenceRunner(configs)\n            dataset = InferenceDataset(configs)\n\n            for i in tqdm(range(len(dataset)), desc=\"Protenix Inference\"):\n                data, atom_array, error_message = dataset[i]\n                target_id = data.get(\"sample_name\", f\"sample_{i}\")\n\n                if target_id not in protenix_queue:\n                    continue\n\n                n_needed, full_seq = protenix_queue[target_id]\n\n                if error_message:\n                    print(f\"  {target_id}: data error — {error_message}\")\n                    protenix_preds[target_id] = None\n                    del data, atom_array, error_message\n                    gc.collect()\n                    continue\n\n                try:\n                    new_cfg = update_inference_configs(configs, data[\"N_token\"].item())\n                    new_cfg.sample_diffusion.N_sample = n_needed\n                    runner.update_model_configs(new_cfg)\n\n                    prediction = runner.predict(data)\n                    raw_coords = prediction[\"coordinate\"]\n                    feat = data[\"input_feature_dict\"]\n\n                    if \"centre_atom_mask\" in feat:\n                        mask = (feat[\"centre_atom_mask\"] == 1).to(raw_coords.device)\n                    elif \"atom_to_tokatom_idx\" in feat:\n                        m11 = (feat[\"atom_to_tokatom_idx\"] == 11).to(raw_coords.device)\n                        m12 = (feat[\"atom_to_tokatom_idx\"] == 12).to(raw_coords.device)\n                        mask = m11 if abs(m11.sum() - len(full_seq)) < abs(m12.sum() - len(full_seq)) else m12\n                    else:\n                        mask = torch.zeros(raw_coords.shape[1], dtype=torch.bool, device=raw_coords.device)\n\n                    coords = raw_coords[:, mask, :].detach().cpu().numpy()\n\n                    # Model collapse check\n                    if coords.shape[1] > 1:\n                        diffs = np.linalg.norm(coords[0, 1:] - coords[0, :-1], axis=-1)\n                        if np.all(diffs < 1e-4):\n                            print(f\"  WARNING: {target_id}: Model collapse detected.\")\n                            coords = np.zeros((coords.shape[0], len(full_seq), 3))\n\n                    # Handle length mismatch\n                    current_len = coords.shape[1]\n                    target_len = len(full_seq)\n                    if current_len != target_len:\n                        padded = np.zeros((coords.shape[0], target_len, 3), dtype=np.float32)\n                        min_len = min(current_len, target_len)\n                        if min_len > 0:\n                            padded[:, :min_len, :] = coords[:, :min_len, :]\n                            # V11+: End-same padding instead of zero padding\n                            if VERSION >= 11 and min_len < target_len:\n                                padded[:, min_len:, :] = coords[:, min_len-1:min_len, :]\n                        coords = padded\n\n                    protenix_preds[target_id] = coords\n\n                except Exception as exc:\n                    print(f\"  {target_id}: Protenix FAILED — {exc}\")\n                    protenix_preds[target_id] = None\n\n                finally:\n                    del data, atom_array\n                    if 'prediction' in locals(): del prediction\n                    if 'raw_coords' in locals(): del raw_coords\n                    torch.cuda.empty_cache()\n                    gc.collect()\n\n        except Exception as e:\n            print(f\"CRITICAL PROTENIX SETUP FAILURE: {e}\")\n            import traceback\n            traceback.print_exc()\n\n    elif protenix_queue and not USE_PROTENIX:\n        print(f\"\\nPHASE 2 skipped (USE_PROTENIX=False).\")\n\n    # ─── PHASE 3: Combine & Output ───\n    print(f\"\\n{'='*60}\")\n    print(\"PHASE 3: Combine TBM + Protenix + de-novo fallback\")\n    print(f\"{'='*60}\")\n\n    all_rows = []\n\n    for _, row in tqdm(test_df.iterrows(), total=len(test_df), desc=\"Merging\"):\n        tid = row[\"target_id\"]\n        seq = row[\"sequence\"]\n\n        combined = list(template_preds.get(tid, []))\n\n        ptx = protenix_preds.get(tid)\n        if ptx is not None:\n            if ptx.ndim == 3:\n                for j in range(ptx.shape[0]):\n                    if len(combined) >= N_SAMPLE:\n                        break\n                    combined.append(ptx[j])\n\n        while len(combined) < N_SAMPLE:\n            seed_val = row.name * 1_000_000 + len(combined) * 1_000\n            dn = generate_rna_structure(seq, seed=seed_val)\n            refined_dn = adaptive_rna_constraints(dn, tid, segments_map, confidence=0.2,\n                                                   seq=seq if VERSION >= 12 else None)\n            combined.append(refined_dn)\n\n        stacked = np.stack(combined[:N_SAMPLE], axis=0)\n        rows = coords_to_rows(tid, seq, stacked)\n        all_rows.extend(rows)\n\n    # ── Save ──\n    sub = pd.DataFrame(all_rows)\n    cols = [\"ID\", \"resname\", \"resid\"] + [\n        f\"{c}_{i}\" for i in range(1, N_SAMPLE + 1) for c in [\"x\", \"y\", \"z\"]\n    ]\n    coord_cols = [c for c in cols if c.startswith((\"x_\", \"y_\", \"z_\"))]\n    sub[coord_cols] = sub[coord_cols].fillna(0.0).clip(-999.999, 9999.999)\n    sub = sub[cols]\n    sub.to_csv(output_csv, index=False)\n    print(f\"\\n✓ Saved submission to {output_csv}  ({len(sub):,} rows)\")\n\n\n# %%  ─── Cell 4: Execute ───","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Main\nRun the prediction pipeline","metadata":{}},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    main()","metadata":{},"outputs":[],"execution_count":null}]}