{"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.12"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":118765,"databundleVersionId":15231210,"isSourceIdPinned":false},{"sourceType":"datasetVersion","sourceId":10880419,"datasetId":6760509,"databundleVersionId":11247150},{"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}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":3984.531263,"end_time":"2026-03-23T03:07:32.92992","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2026-03-23T02:01:08.398657","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"5cf539e6","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":{"execution":{"iopub.execute_input":"2026-03-23T02:01:10.977806Z","iopub.status.busy":"2026-03-23T02:01:10.977501Z","iopub.status.idle":"2026-03-23T02:01:11.245917Z","shell.execute_reply":"2026-03-23T02:01:11.244941Z"},"papermill":{"duration":0.273953,"end_time":"2026-03-23T02:01:11.24735","exception":false,"start_time":"2026-03-23T02:01:10.973397","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"ffa23784","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":{"execution":{"iopub.execute_input":"2026-03-23T02:01:11.253803Z","iopub.status.busy":"2026-03-23T02:01:11.253547Z","iopub.status.idle":"2026-03-23T02:02:08.483094Z","shell.execute_reply":"2026-03-23T02:02:08.481903Z"},"papermill":{"duration":57.234486,"end_time":"2026-03-23T02:02:08.484849","exception":false,"start_time":"2026-03-23T02:01:11.250363","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"e22a9631","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":{"execution":{"iopub.execute_input":"2026-03-23T02:02:08.497918Z","iopub.status.busy":"2026-03-23T02:02:08.497529Z","iopub.status.idle":"2026-03-23T02:02:09.22113Z","shell.execute_reply":"2026-03-23T02:02:09.220384Z"},"papermill":{"duration":0.731092,"end_time":"2026-03-23T02:02:09.222676","exception":false,"start_time":"2026-03-23T02:02:08.491584","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"98cf97a1","cell_type":"code","source":"print(\"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\"Loaded {len(train_seqs)} training sequences\")\nprint(f\"Loaded {len(validation_seqs)} validation sequences\")\nprint(f\"Loaded {len(test_seqs)} test sequences\")\n\nSHOW_VALIDATION = False\nSHOW_ALIGNMENT_DETAILS = True\nMAKE_SUBMISSION = True\nUSE_PROTENIX = True\n\nMIN_SIMILARITY = 0.15\nMIN_PERCENT_IDENTITY = 35\n\nDEBUG = False\nif DEBUG:\n    MIN_SIMILARITY = 0.3","metadata":{"execution":{"iopub.execute_input":"2026-03-23T02:02:09.234657Z","iopub.status.busy":"2026-03-23T02:02:09.234284Z","iopub.status.idle":"2026-03-23T02:02:19.089386Z","shell.execute_reply":"2026-03-23T02:02:19.088319Z"},"papermill":{"duration":9.862581,"end_time":"2026-03-23T02:02:19.090881","exception":false,"start_time":"2026-03-23T02:02:09.2283","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"2b48ac00","cell_type":"code","source":"def make_aligner():\n    al = PairwiseAligner()\n    al.mode = 'global'\n    al.match_score = 3.0\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 = -5\n    al.query_left_extend_gap_score = -0.5\n    al.query_right_open_gap_score = -5\n    al.query_right_extend_gap_score = -0.5\n    al.target_left_open_gap_score = -5\n    al.target_left_extend_gap_score = -0.5\n    al.target_right_open_gap_score = -5\n    al.target_right_extend_gap_score = -0.5\n    return al\n\n_aligner = make_aligner()\n\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\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\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\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\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\n\ndef compute_sequence_identity(seq1, seq2):\n    alignment = next(iter(_aligner.align(seq1, seq2)))\n    identical = 0\n    total_aligned = 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            total_aligned += 1\n            if seq1[q_pos] == seq2[t_pos]:\n                identical += 1\n    return 100 * identical / max(len(seq1), 1)\n\n\ndef find_similar_sequences(query_seq, train_seqs_df, train_coords_dict, top_n=10):\n    similar_seqs = []\n\n    for _, row in train_seqs_df.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        raw_score = _aligner.score(query_seq, train_seq)\n        max_score = 3.0 * min(len(query_seq), len(train_seq))\n        normalized_score = raw_score / max(max_score, 1)\n        \n        similar_seqs.append((target_id, train_seq, normalized_score, train_coords_dict[target_id]))\n\n    similar_seqs.sort(key=lambda x: x[2], reverse=True)\n    return similar_seqs[:top_n]\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\n\ndef find_similar_sequences_detailed(query_seq, train_seqs_df, train_coords_dict,\n                                    temporal_cutoff=None, top_n=10):\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\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        alignment = next(iter(_aligner.align(query_seq, train_seq)))\n        raw_score = alignment.score\n        max_score = 3.0 * min(len(query_seq), len(train_seq))\n        normalized_score = raw_score / max(max_score, 1)\n\n        identical = 0\n        total_aligned = 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                total_aligned += 1\n                if query_seq[q_pos] == train_seq[t_pos]:\n                    identical += 1\n        \n        percent_identity = 100 * identical / max(len(query_seq), 1)\n\n        aligned_query, aligned_template = _build_aligned_strings(\n            query_seq, train_seq, alignment\n        )\n\n        similar_seqs.append((\n            target_id, train_seq, 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], x[4]), reverse=True)\n    return similar_seqs[:top_n]\n\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                if i == prev_v + 1:\n                    new_coords[i] = new_coords[prev_v] + [5.95, 0, 0]\n                else:\n                    vec = np.array([5.95, 0, 0])\n                    new_coords[i] = new_coords[prev_v] + vec * (i - prev_v)\n            elif next_v >= 0:\n                if i == next_v - 1:\n                    new_coords[i] = new_coords[next_v] - [5.95, 0, 0]\n                else:\n                    vec = np.array([5.95, 0, 0])\n                    new_coords[i] = new_coords[next_v] - vec * (next_v - i)\n            else:\n                new_coords[i] = [i * 5.95, 0, 0]\n\n    return np.nan_to_num(new_coords)\n\n\ndef adaptive_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    strength = 0.85 * (1.0 - min(confidence, 0.95))\n    strength = max(strength, 0.03)\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            d = X[1:] - X[:-1]\n            dist = np.linalg.norm(d, axis=1) + 1e-6\n            target = 5.95\n            scale = (target - dist) / dist\n            adj = (d * scale[:, None]) * (0.25 * strength)\n            X[:-1] -= adj\n            X[1:] += adj\n\n            if L > 2:\n                d2 = X[2:] - X[:-2]\n                dist2 = np.linalg.norm(d2, axis=1) + 1e-6\n                target2 = 10.2\n                scale2 = (target2 - dist2) / dist2\n                adj2 = (d2 * scale2[:, None]) * (0.12 * strength)\n                X[:-2] -= adj2\n                X[2:] += adj2\n\n            if L > 2:\n                lap = 0.5 * (X[:-2] + X[2:]) - X[1:-1]\n                X[1:-1] += (0.08 * strength) * lap\n\n            if L >= 25:\n                k = min(L, 180) if L > 220 else L\n                if k < L:\n                    idx = np.linspace(0, L - 1, k).astype(int)\n                else:\n                    idx = np.arange(L)\n\n                P = X[idx]\n                diff = P[:, None, :] - P[None, :, :]\n                distm = np.linalg.norm(diff, axis=2) + 1e-6\n                sep = np.abs(idx[:, None] - idx[None, :])\n\n                mask = (sep > 2) & (distm < 3.3)\n                if np.any(mask):\n                    force = (3.3 - 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\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\n\ndef apply_hinge(coords, seg, rng, max_angle_deg=20):\n    s, e = seg\n    L = e - s\n    if L < 30:\n        return coords\n    pivot = s + int(rng.integers(10, L - 10))\n    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\n\ndef jitter_chains(coords, segments, rng, max_angle_deg=10, max_trans=1.2):\n    X = coords.copy()\n    global_center = 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) - global_center\n    return X\n\n\ndef smooth_wiggle(coords, segments, rng, amp=0.6):\n    X = coords.copy()\n    for (s, e) in segments:\n        L = e - s\n        if L < 20:\n            continue\n        n_ctrl = 6\n        ctrl_x = np.linspace(0, L - 1, n_ctrl)\n        ctrl_disp = rng.normal(0, amp, size=(n_ctrl, 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\ndef predict_rna_structures(row, train_seqs_df, train_coords_dict, segments_map, n_predictions=5):\n    tid = row['target_id']\n    seq = row['sequence']\n    segments = segments_map.get(tid, [(0, len(seq))])\n\n    cands = find_similar_sequences(\n        query_seq=seq, train_seqs_df=train_seqs_df,\n        train_coords_dict=train_coords_dict, top_n=40\n    )\n\n    predictions = []\n    used = set()\n\n    for i in range(n_predictions):\n        seed = (abs(hash(tid)) + i * 10007) % (2**32)\n        rng = np.random.default_rng(seed)\n\n        if not cands:\n            coords = np.zeros((len(seq), 3), dtype=float)\n            for (s, e) in segments:\n                for j in range(s + 1, e):\n                    coords[j] = coords[j - 1] + [5.95, 0, 0]\n            predictions.append(coords)\n            continue\n\n        if i == 0:\n            t_id, t_seq, sim, t_coords = cands[0]\n        else:\n            K = min(15, len(cands))\n            sims = np.array([cands[k][2] for k in range(K)], float)\n            w = np.exp((sims - sims.max()) / 0.10)\n            for k in range(K):\n                if cands[k][0] in used:\n                    w[k] *= 0.05\n            w = w / (w.sum() + 1e-12)\n            k = int(rng.choice(np.arange(K), p=w))\n            t_id, t_seq, sim, t_coords = cands[k]\n\n        used.add(t_id)\n\n        adapted = adapt_template_to_query(\n            query_seq=seq, template_seq=t_seq, template_coords=t_coords\n        )\n\n        if i == 0:\n            X = adapted\n        elif i == 1:\n            X = adapted + rng.normal(0, max(0.01, (0.35 - sim) * 0.08), adapted.shape)\n        elif i == 2:\n            longest = max(segments, key=lambda se: se[1] - se[0])\n            X = apply_hinge(adapted, longest, rng, max_angle_deg=18)\n        elif i == 3:\n            X = jitter_chains(adapted, segments, rng, max_angle_deg=8, max_trans=0.8)\n        else:\n            X = smooth_wiggle(adapted, segments, rng, amp=0.5)\n\n        refined = adaptive_rna_constraints(X, tid, segments_map, confidence=sim, passes=3)\n        predictions.append(refined)\n\n    return predictions\n\n\ndef generate_rna_structure(sequence, seed=None):\n    \"\"\"Nussinov+A-form geometry — much better than simple helix.\"\"\"\n    _p = {(\"A\",\"U\"),(\"U\",\"A\"),(\"G\",\"C\"),(\"C\",\"G\"),(\"G\",\"U\"),(\"U\",\"G\")}\n    AR=2.81; AT=np.radians(32.7); ARAD=9.0; SS=5.9\n    rng=np.random.default_rng(seed if seed is not None else 42)\n    seq=sequence.upper().replace(\"T\",\"U\"); n=len(seq)\n    dp=[[0]*n for _ in range(n)]; bt=[[None]*n for _ in range(n)]\n    for s in range(4,n):\n        for i in range(n-s):\n            j=i+s; bst=dp[i+1][j]; src=(\"L\",i+1,j)\n            if dp[i][j-1]>bst: bst=dp[i][j-1]; src=(\"R\",i,j-1)\n            if (seq[i],seq[j]) in _p:\n                v=dp[i+1][j-1]+1\n                if v>bst: bst=v; src=(\"P\",i+1,j-1)\n            for k in range(i+1,j):\n                v=dp[i][k]+dp[k+1][j]\n                if v>bst: bst=v; src=(\"B\",k,0)\n            dp[i][j]=bst; bt[i][j]=src\n    pairs=[]; stack=[(0,n-1)]\n    while stack:\n        i,j=stack.pop()\n        if i>=j or bt[i][j] is None: continue\n        t=bt[i][j][0]\n        if t in (\"L\",\"R\"): stack.append((bt[i][j][1],bt[i][j][2]))\n        elif t==\"P\":\n            pairs.append((i,j)); a,b=bt[i][j][1],bt[i][j][2]\n            if a<=b: stack.append((a,b))\n        elif t==\"B\": k=bt[i][j][1]; stack.extend([(i,k),(k+1,j)])\n    paired={i:j for i,j in pairs}; paired.update({j:i for i,j in pairs})\n    xyz=np.zeros((n,3),dtype=float); done=np.zeros(n,dtype=bool)\n    cur=np.zeros(3); hd=np.array([0.,0.,1.]); i=0\n    while i<n:\n        if i in paired and paired[i]>i:\n            j=paired[i]; a0=rng.uniform(0,2*np.pi)\n            for k in range(j-i+1):\n                a=a0+k*AT\n                xyz[i+k]=cur+[ARAD*np.cos(a),ARAD*np.sin(a),k*AR]; done[i+k]=True\n                if not done[j-k]:\n                    xyz[j-k]=cur+[ARAD*np.cos(a+np.pi),ARAD*np.sin(a+np.pi),k*AR]\n                    done[j-k]=True\n            cur=xyz[j]+hd*3.; i=j+1\n        else:\n            end=i+1\n            while end<n and (end not in paired or paired[end]<i): end+=1\n            for k in range(end-i):\n                if not done[i+k]: xyz[i+k]=cur+k*SS*hd; done[i+k]=True\n            if end>i: cur=xyz[end-1]+hd*3.\n            i=end\n    for idx in range(n):\n        if not done[idx]: xyz[idx]=cur+idx*hd*SS\n    return xyz","metadata":{"execution":{"iopub.execute_input":"2026-03-23T02:02:19.104832Z","iopub.status.busy":"2026-03-23T02:02:19.104452Z","iopub.status.idle":"2026-03-23T02:02:19.150643Z","shell.execute_reply":"2026-03-23T02:02:19.150024Z"},"papermill":{"duration":0.054776,"end_time":"2026-03-23T02:02:19.151862","exception":false,"start_time":"2026-03-23T02:02:19.097086","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"f302eff3","cell_type":"code","source":"print(\"Processing coordinates...\")\ntrain_coords_dict = process_labels(train_labels)\n\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\nprint(f\"Processed {len(train_coords_dict)} training structures\")\nprint(f\"Processed {len(combined_coords_dict)} combined (train+val) structures\")\n","metadata":{"execution":{"iopub.execute_input":"2026-03-23T02:02:19.163167Z","iopub.status.busy":"2026-03-23T02:02:19.162923Z","iopub.status.idle":"2026-03-23T02:03:07.084772Z","shell.execute_reply":"2026-03-23T02:03:07.083826Z"},"papermill":{"duration":47.933738,"end_time":"2026-03-23T02:03:07.090933","exception":false,"start_time":"2026-03-23T02:02:19.157195","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"7ff91fd3","cell_type":"code","source":"validation_segments_map, _ = build_segments_map(validation_seqs)\ntest_segments_map, _ = build_segments_map(test_seqs)\n\nprint(f\"Built segment maps: {len(validation_segments_map)} validation, {len(test_segments_map)} test\")","metadata":{"execution":{"iopub.execute_input":"2026-03-23T02:03:07.102257Z","iopub.status.busy":"2026-03-23T02:03:07.101994Z","iopub.status.idle":"2026-03-23T02:03:07.109858Z","shell.execute_reply":"2026-03-23T02:03:07.109026Z"},"papermill":{"duration":0.014739,"end_time":"2026-03-23T02:03:07.11102","exception":false,"start_time":"2026-03-23T02:03:07.096281","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"d6750302","cell_type":"code","source":"from biotite.structure.io.pdbx import CIFFile, get_structure\n\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\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        print(f\"    Sequence too long ({len(sequence)} > {max_length}), truncating, no MSA\")\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\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\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@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\ntemplate_info_dict = {}\nprediction_metadata_dict = {}\n\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\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\ndef predict_with_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(\"\\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 + 5\n    )\n\n    if similar_seqs:\n        for i, (tmpl_id, tmpl_seq, 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):\n                if SHOW_ALIGNMENT_DETAILS:\n                    print(f\"  Template {i+1}: {tmpl_id} SKIPPED \"\n                          f\"(sim={similarity:.3f}, id={pct_id:.1f}%) - below threshold\")\n                break\n\n            if USE_PROTENIX and len(aligned_q) < 100 and i == 4:\n                print(f\"  Sequence length {len(aligned_q)} short, leaving 1 space for Protenix\")\n                break\n\n            record_template_info(target_id, tmpl_id, similarity, pct_id)\n            record_prediction_metadata(target_id, pred_num, 'template',\n                                       tmpl_id, similarity, pct_id)\n\n            if SHOW_ALIGNMENT_DETAILS:\n                print(f\"  Template {i+1}: {tmpl_id} (sim={similarity:.3f}, id={pct_id:.1f}%)\")\n\n            adapted = adapt_template_to_query(sequence, tmpl_seq, tmpl_coords)\n            refined = adaptive_rna_constraints(adapted, target_id, segments_map,\n                                               confidence=similarity, passes=3)\n            predictions.append(refined)\n            pred_num += 1\n\n            if len(predictions) >= n_predictions:\n                break\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} from templates, {n_needed} slots for Protenix\")\n    else:\n        print(f\"  -> All {n_predictions} predictions from templates\")\n\n    return predictions, n_needed, pred_num\n\n\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    \"\"\"\n    Enhanced generate_predictions_batch:\n    - Runs Protenix with 3 seeds x 3 samples for ALL targets\n    - Picks best 5 from TBM + Protenix pool ranked by score (pLDDT / similarity)\n    - Everything else identical to original (same column order, same row format)\n    \"\"\"\n    start_time = time.time()\n    total_targets = len(sequences_df)\n\n    print(f\"\\n{'='*70}\")\n    print(f\"Predicting {total_targets} {dataset_name} sequences\")\n    print(f\"{'='*70}\")\n\n    segments_map, _ = build_segments_map(sequences_df)\n\n    # ---- Phase 1: Template predictions (identical to original) ----\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        temporal_cutoff = 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=temporal_cutoff\n        )\n        template_predictions[target_id] = preds\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\")\n\n    # ---- Phase 2: Protenix multi-seed for ALL targets ----\n    protenix_candidates = {}   # tid -> [(plddt, coords), ...]\n\n    if USE_PROTENIX:\n        SEEDS = (101, 102, 103)\n        N_PER_SEED = 3\n\n        print(f\"\\nPHASE 2: Protenix seeds={SEEDS} x {N_PER_SEED} samples each target\")\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 idx, (target_id, (n_needed, next_pred, sequence)) in enumerate(protenix_queue.items()):\n                print(f\"\\n  [{idx+1}/{len(protenix_queue)}] {target_id} ({len(sequence)} nt)\")\n\n                try:\n                    prepare_protenix_json(target_id, sequence,\n                                         protenix_output_path, input_path)\n                    all_cands = []\n\n                    for seed in SEEDS:\n                        print(f\"    seed={seed}...\", end=\" \", flush=True)\n                        t0 = time.time()\n                        try:\n                            run_protenix_inference(\n                                target_id, sequence, protenix_output_path, input_path,\n                                seed=seed, n_cycle=12, n_sample=N_PER_SEED, n_step=250\n                            )\n                            # Collect CIFs with pLDDT scores\n                            px_base = protenix_output_path / target_id\n                            for si in range(N_PER_SEED):\n                                found_cif = None\n                                # Try all known path patterns\n                                for cif_try in [\n                                    px_base / target_id / f\"seed_{seed}\" / \"predictions\" / f\"{target_id}_seed_{seed}_sample_{si}.cif\",\n                                    px_base / f\"seed_{seed}\" / \"predictions\" / f\"{target_id}_seed_{seed}_sample_{si}.cif\",\n                                    px_base / str(seed) / \"predictions\" / f\"{target_id}_seed_{seed}_sample_{si}.cif\",\n                                ]:\n                                    if cif_try.exists():\n                                        found_cif = cif_try\n                                        break\n                                if found_cif is None:\n                                    hits = sorted(px_base.glob(f\"**/*sample_{si}*.cif\"))\n                                    if hits:\n                                        found_cif = hits[0]\n                                if found_cif is not None:\n                                    try:\n                                        pred_df = extract_c1_atoms(str(found_cif))\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                                        plddt = 0.5\n                                        for conf_path in (\n                                            sorted(found_cif.parent.glob(f\"*confidence*sample_{si}*.json\")) +\n                                            sorted(found_cif.parent.glob(f\"*summary*sample_{si}*.json\"))\n                                        ):\n                                            try:\n                                                with open(conf_path) as cf:\n                                                    d = json.load(cf)\n                                                plddt = float(d.get(\"plddt\", d.get(\"mean_plddt\", 0.5)))\n                                                break\n                                            except Exception:\n                                                pass\n                                        all_cands.append((plddt, coords))\n                                    except Exception as ce:\n                                        print(f\"      parse error: {ce}\")\n                        except Exception as se:\n                            print(f\"      seed={seed} failed: {se}\")\n                        print(f\"{(time.time()-t0)/60:.1f}min\")\n\n                    all_cands.sort(key=lambda x: -x[0])\n                    protenix_candidates[target_id] = all_cands\n                    if all_cands:\n                        print(f\"    {len(all_cands)} candidates  best_pLDDT={all_cands[0][0]:.3f}\")\n\n                except Exception as e:\n                    print(f\"    FAILED: {e}\")\n                    protenix_candidates[target_id] = []\n\n        ptx_time = time.time() - start_time - template_time\n        print(f\"\\nPhase 2 done: {ptx_time/60:.1f} min\")\n\n    # ---- Phase 3: Pick best 5 from all sources ----\n    print(f\"\\nPHASE 3: Selecting best 5 per target\")\n\n    all_rows = []\n\n    for _, row in sequences_df.iterrows():\n        target_id = row['target_id']\n        sequence = row['sequence']\n\n        # Build scored pool: TBM by similarity, Protenix by pLDDT\n        all_cands = []\n        for sim, coords in template_predictions.get(target_id, []):\n            all_cands.append((float(sim), coords, 'template'))\n        for plddt, coords in protenix_candidates.get(target_id, []):\n            all_cands.append((float(plddt), coords, 'protenix'))\n        all_cands.sort(key=lambda x: -x[0])\n        top5 = all_cands[:5]\n\n        # Pad with de novo geometry if needed\n        n_denovo = 0\n        while len(top5) < 5:\n            seed_val = hash(target_id) % 10000 + len(top5) * 1000\n            de_novo = generate_rna_structure(sequence, seed=seed_val)\n            refined = adaptive_rna_constraints(de_novo, target_id, segments_map,\n                                               confidence=0.15, passes=3)\n            top5.append((0.05, refined, 'de_novo'))\n            n_denovo += 1\n\n        sources = [s for _,_,s in top5]\n        scores  = [f\"{sc:.3f}\" for sc,_,_ in top5]\n        print(f\"  {target_id}: {sources} scores={scores}\" +\n              (f\" [{n_denovo} de_novo]\" if n_denovo else \"\"))\n\n        # Build rows — EXACT same format as original\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}'] = top5[i][1][j][0]\n                pred_row[f'y_{i+1}'] = top5[i][1][j][1]\n                pred_row[f'z_{i+1}'] = top5[i][1][j][2]\n            all_rows.append(pred_row)\n\n    submission_df = pd.DataFrame(all_rows)\n    # EXACT same column order as original\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(\n        1 for tid in sequences_df['target_id']\n        if not protenix_candidates.get(tid)\n    )\n    n_with_protenix = total_targets - n_template_only\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: {n_with_protenix}\")\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","metadata":{"execution":{"iopub.execute_input":"2026-03-23T02:03:07.122601Z","iopub.status.busy":"2026-03-23T02:03:07.122389Z","iopub.status.idle":"2026-03-23T02:03:07.839767Z","shell.execute_reply":"2026-03-23T02:03:07.838988Z"},"papermill":{"duration":0.725028,"end_time":"2026-03-23T02:03:07.841356","exception":false,"start_time":"2026-03-23T02:03:07.116328","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"a4fa9021","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\")\ntest_predictions.head(10)\n","metadata":{"execution":{"iopub.execute_input":"2026-03-23T02:03:07.853506Z","iopub.status.busy":"2026-03-23T02:03:07.853093Z","iopub.status.idle":"2026-03-23T03:07:28.999249Z","shell.execute_reply":"2026-03-23T03:07:28.998249Z"},"papermill":{"duration":3861.153533,"end_time":"2026-03-23T03:07:29.000666","exception":false,"start_time":"2026-03-23T02:03:07.847133","status":"completed"},"tags":[]},"outputs":[],"execution_count":null}]}