{"metadata":{"kernelspec":{"display_name":"Python 3","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.12.12"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":118765,"databundleVersionId":15231210,"sourceType":"competition"},{"sourceId":14591975,"sourceType":"datasetVersion","datasetId":9320850},{"sourceId":14604295,"sourceType":"datasetVersion","datasetId":9328538}],"dockerImageVersionId":31260,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":2354.045049,"end_time":"2026-01-13T06:54:44.524542","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2026-01-13T06:15:30.479493","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install --no-index /kaggle/input/biopython-wheel/biopython-1.86-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-12T12:26:37.327727Z","iopub.execute_input":"2026-02-12T12:26:37.328282Z","iopub.status.idle":"2026-02-12T12:26:42.637608Z","shell.execute_reply.started":"2026-02-12T12:26:37.328227Z","shell.execute_reply":"2026-02-12T12:26:42.636944Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Kaggle Notebook — Multi-template blending + MC + template-contacts + anchor springs\n# Robust to train_labels containing only x_1,y_1,z_1 (no x_2..x_5)\n# Produces submission.csv for Stanford RNA 3D Folding 2\n\nimport os, sys, time, warnings\nimport numpy as np\nimport pandas as pd\nfrom collections import defaultdict\n\nwarnings.filterwarnings(\"ignore\")\n\nDATA_PATH = \"/kaggle/input/stanford-rna-3d-folding-2/\"\n\ntrain_seqs   = pd.read_csv(DATA_PATH + \"train_sequences.csv\")\ntest_seqs    = pd.read_csv(DATA_PATH + \"test_sequences.csv\")\ntrain_labels = pd.read_csv(DATA_PATH + \"train_labels.csv\")\n\n# -----------------------------\n# FASTA + stoichiometry parsing\n# -----------------------------\nsys.path.append(os.path.join(DATA_PATH, \"extra\"))\n\ntry:\n    import typing as _typing\n    import builtins as _builtins\n    _builtins.Dict  = getattr(_typing, \"Dict\")\n    _builtins.Tuple = getattr(_typing, \"Tuple\")\n    _builtins.List  = getattr(_typing, \"List\")\n    from parse_fasta_py import parse_fasta as _parse_fasta_raw\n\n    def parse_fasta(fasta_content: str):\n        d = _parse_fasta_raw(fasta_content)\n        out = {}\n        for k, v in d.items():\n            out[k] = v[0] if isinstance(v, tuple) else v\n        return out\n\nexcept Exception:\n    def 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                header = line[1:]\n                cur = header.split()[0]\n                seq_parts = []\n            else:\n                seq_parts.append(line.replace(\" \", \"\"))\n        if cur is not None:\n            out[cur] = \"\".join(seq_parts)\n        return out\n\ndef parse_stoichiometry(stoich: str):\n    if pd.isna(stoich) or str(stoich).strip() == \"\":\n        return []\n    out = []\n    for part in str(stoich).split(';'):\n        ch, cnt = part.split(':')\n        out.append((ch.strip(), int(cnt)))\n    return out\n\ndef get_chain_segments(row):\n    seq = row[\"sequence\"]\n    stoich = row.get(\"stoichiometry\", \"\")\n    all_seq = row.get(\"all_sequences\", \"\")\n\n    if pd.isna(stoich) or pd.isna(all_seq) or str(stoich).strip()==\"\" or str(all_seq).strip()==\"\":\n        return [(0, len(seq))]\n\n    try:\n        chain_dict = parse_fasta(all_seq)\n        order = parse_stoichiometry(stoich)\n        segs = []\n        pos = 0\n        for ch, cnt in order:\n            base = chain_dict.get(ch)\n            if base is None:\n                return [(0, len(seq))]\n            for _ in range(cnt):\n                L = len(base)\n                segs.append((pos, pos + L))\n                pos += L\n        if pos != len(seq):\n            return [(0, len(seq))]\n        return segs\n    except Exception:\n        return [(0, len(seq))]\n\ndef build_segments_map(df):\n    seg_map = {}\n    for _, r in df.iterrows():\n        seg_map[r[\"target_id\"]] = get_chain_segments(r)\n    return seg_map\n\ntrain_segs_map = build_segments_map(train_seqs)\ntest_segs_map  = build_segments_map(test_seqs)\n\n# -----------------------------\n# Detect how many coord sets exist in train_labels (often only 1)\n# -----------------------------\ndef detect_coord_sets(labels_df):\n    # looks for x_1,y_1,z_1 ... x_k,y_k,z_k\n    k = 1\n    while all(col in labels_df.columns for col in [f\"x_{k}\", f\"y_{k}\", f\"z_{k}\"]):\n        k += 1\n    return k - 1\n\nN_COORD_SETS = detect_coord_sets(train_labels)\nif N_COORD_SETS < 1:\n    raise RuntimeError(\"train_labels has no x_1,y_1,z_1 columns. Unexpected schema.\")\n\nprint(f\"Detected {N_COORD_SETS} coordinate set(s) in train_labels (x_1..x_{N_COORD_SETS}).\")\n\n# -----------------------------\n# Build template coordinate dict (stores (K, L, 3) even if K=1)\n# -----------------------------\ndef process_labels_k(labels_df, K):\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        g = group.sort_values(\"resid\")\n        coordsK = []\n        for k in range(1, K + 1):\n            coordsK.append(g[[f\"x_{k}\", f\"y_{k}\", f\"z_{k}\"]].values.astype(np.float32))\n        coords_dict[id_prefix] = np.stack(coordsK, axis=0)  # (K, L, 3)\n    return coords_dict\n\ntrain_coordsK_dict = process_labels_k(train_labels, N_COORD_SETS)\n\n# -----------------------------\n# Alignment (CPU; fast score)\n# -----------------------------\nfrom Bio.Align import PairwiseAligner\n\naligner = PairwiseAligner()\naligner.mode = \"global\"\naligner.match_score = 2\naligner.mismatch_score = -1.5\n\naligner.open_gap_score   = -8\naligner.extend_gap_score = -0.4\n\naligner.query_left_open_gap_score  = -8\naligner.query_left_extend_gap_score = -0.4\naligner.query_right_open_gap_score = -8\naligner.query_right_extend_gap_score = -0.4\naligner.target_left_open_gap_score = -8\naligner.target_left_extend_gap_score = -0.4\naligner.target_right_open_gap_score = -8\naligner.target_right_extend_gap_score = -0.4\n\ndef find_similar_sequences(query_seq, train_seqs_df, train_coordsK_dict, top_n=30):\n    sims = []\n    qL = len(query_seq)\n    for _, row in train_seqs_df.iterrows():\n        tid = row[\"target_id\"]\n        if tid not in train_coordsK_dict:\n            continue\n        tseq = row[\"sequence\"]\n        tL = len(tseq)\n        if abs(tL - qL) / max(tL, qL) > 0.30:\n            continue\n\n        raw_score = aligner.score(query_seq, tseq)\n        norm = raw_score / (2.0 * min(qL, tL))\n        sims.append((tid, tseq, float(norm), train_coordsK_dict[tid]))  # coords: (K,Lt,3)\n\n    sims.sort(key=lambda x: x[2], reverse=True)\n    return sims[:top_n]\n\n# -----------------------------\n# GPU helpers (torch)\n# -----------------------------\nimport torch\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\ntorch.set_grad_enabled(False)\n\ndef kabsch_align(P, Q):\n    Pc = P.mean(dim=0, keepdim=True)\n    Qc = Q.mean(dim=0, keepdim=True)\n    P0 = P - Pc\n    Q0 = Q - Qc\n    H = Q0.T @ P0\n    U, S, Vt = torch.linalg.svd(H)\n    R = Vt.T @ U.T\n    if torch.linalg.det(R) < 0:\n        Vt[-1, :] *= -1\n        R = Vt.T @ U.T\n    Q_aligned = (Q - Qc) @ R + Pc\n    return Q_aligned\n\n# -----------------------------\n# Gap fill + anchor mask\n# -----------------------------\ndef adapt_template_to_query_dirfill(query_seq, template_seq, template_coords, segments):\n    alignment = next(iter(aligner.align(query_seq, template_seq)))\n    Lq = len(query_seq)\n    new_coords = np.full((Lq, 3), np.nan, dtype=np.float32)\n    anchor = np.zeros((Lq,), dtype=bool)\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            anchor[q_start:q_end] = True\n\n    step = 5.95\n    for (s, e) in segments:\n        X = new_coords[s:e]\n        L = e - s\n        if L <= 1:\n            continue\n\n        for i in range(L):\n            if not np.isnan(X[i, 0]):\n                continue\n            prev = i - 1\n            while prev >= 0 and np.isnan(X[prev, 0]):\n                prev -= 1\n            nxt = i + 1\n            while nxt < L and np.isnan(X[nxt, 0]):\n                nxt += 1\n\n            if prev >= 0 and nxt < L:\n                w = (i - prev) / (nxt - prev)\n                X[i] = (1 - w) * X[prev] + w * X[nxt]\n            elif prev >= 0:\n                if prev >= 1 and not np.isnan(X[prev - 1, 0]):\n                    v = X[prev] - X[prev - 1]\n                    v = v / (np.linalg.norm(v) + 1e-6)\n                else:\n                    v = np.array([1, 0, 0], dtype=np.float32)\n                X[i] = X[prev] + v * step\n            elif nxt < L:\n                if nxt + 1 < L and not np.isnan(X[nxt + 1, 0]):\n                    v = X[nxt] - X[nxt + 1]\n                    v = v / (np.linalg.norm(v) + 1e-6)\n                else:\n                    v = np.array([1, 0, 0], dtype=np.float32)\n                X[i] = X[nxt] + v * step\n            else:\n                X[i] = np.array([i * step, 0, 0], dtype=np.float32)\n\n        new_coords[s:e] = X\n\n    return np.nan_to_num(new_coords).astype(np.float32), anchor\n\n# -----------------------------\n# Template contacts\n# -----------------------------\ndef build_contacts_from_template(X, cutoff=12.0, min_sep=5, max_pairs=600):\n    L = X.shape[0]\n    step = 1 if L <= 250 else 2 if L <= 500 else 3\n    idx = np.arange(0, L, step, dtype=int)\n    P = X[idx]\n\n    diff = P[:, None, :] - P[None, :, :]\n    dist = np.linalg.norm(diff, axis=2)\n    sep = np.abs(idx[:, None] - idx[None, :])\n\n    mask = (sep >= min_sep) & (dist < cutoff)\n    ii, jj = np.where(mask)\n    pairs = []\n    for a, b in zip(ii, jj):\n        if a < b:\n            pairs.append((int(idx[a]), int(idx[b]), float(dist[a, b])))\n\n    pairs.sort(key=lambda t: t[2])\n    return pairs[:max_pairs]\n\n# -----------------------------\n# Refinement (numpy) + anchor springs\n# -----------------------------\ndef adaptive_rna_constraints(coordinates, segments, confidence=1.0, passes=2,\n                             anchor_mask=None, anchor_xyz=None):\n    coords = coordinates.copy()\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            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            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            X[:-2] -= adj2\n            X[2:]  += adj2\n\n            lap = 0.5 * (X[:-2] + X[2:]) - X[1:-1]\n            X[1:-1] += (0.05 * strength) * lap\n\n            if L >= 25:\n                k = min(L, 140)\n                idx = np.linspace(0, L-1, k).astype(int) if k < L else np.arange(L)\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.012 * strength) * vec\n\n            if anchor_mask is not None and anchor_xyz is not None:\n                m = anchor_mask[s:e]\n                if np.any(m):\n                    lam = 0.10 + 0.15 * min(float(confidence), 1.0)\n                    X[m] = (1 - lam) * X[m] + lam * anchor_xyz[s:e][m]\n\n            coords[s:e] = X\n\n    return coords\n\n# -----------------------------\n# Energy (torch) + contacts\n# -----------------------------\ndef energy_torch(X, segments, contacts=None):\n    e = torch.zeros((), device=X.device)\n\n    for (s, eidx) in segments:\n        Y = X[s:eidx]\n        L = Y.shape[0]\n        if L < 4:\n            continue\n\n        d1 = Y[1:] - Y[:-1]\n        dist1 = torch.linalg.norm(d1, dim=1)\n        e = e + torch.mean((dist1 - 5.95) ** 2)\n\n        d2 = Y[2:] - Y[:-2]\n        dist2 = torch.linalg.norm(d2, dim=1)\n        e = e + 0.4 * torch.mean((dist2 - 10.2) ** 2)\n\n        if L >= 25:\n            k = min(L, 120)\n            idx = torch.linspace(0, L-1, k, device=X.device).long()\n            P = Y[idx]\n            diff = P[:, None, :] - P[None, :, :]\n            distm = torch.linalg.norm(diff, dim=2) + 1e-6\n            sep = (idx[:, None] - idx[None, :]).abs()\n            mask = (sep > 2) & (distm < 3.2)\n            if mask.any():\n                e = e + 0.05 * torch.mean((3.2 - distm[mask]) ** 2)\n\n    if contacts:\n        ii = torch.tensor([c[0] for c in contacts], device=X.device, dtype=torch.long)\n        jj = torch.tensor([c[1] for c in contacts], device=X.device, dtype=torch.long)\n        d0 = torch.tensor([c[2] for c in contacts], device=X.device, dtype=torch.float32)\n        dij = torch.linalg.norm(X[jj] - X[ii], dim=1)\n        e = e + 0.15 * torch.mean((dij - d0) ** 2)\n\n    return e\n\n# -----------------------------\n# MC blend search\n# -----------------------------\ndef mc_blend_search(adapted_list, sims, segments, seed=0, steps=120, start=\"softmax\",\n                    anchor_mask=None, anchor_xyz=None, contacts=None, noise_sigma=0.0):\n    rng = np.random.default_rng(seed)\n    M = len(adapted_list)\n    T = np.stack(adapted_list, axis=0)  # (M,L,3)\n\n    sims_np = np.array(sims, dtype=np.float32)\n    if start == \"softmax\":\n        w = np.exp((sims_np - sims_np.max()) / 0.08)\n        w = w / (w.sum() + 1e-12)\n    else:\n        w = np.ones(M, dtype=np.float32) / M\n\n    def blend(w):\n        X = (w[:, None, None] * T).sum(axis=0).astype(np.float32)\n        if noise_sigma > 0:\n            X = X + rng.normal(0.0, noise_sigma, size=X.shape).astype(np.float32)\n        return X\n\n    conf = float(np.max(sims_np)) if len(sims_np) else 0.0\n\n    X = adaptive_rna_constraints(blend(w), segments, confidence=conf, passes=1,\n                                anchor_mask=anchor_mask, anchor_xyz=anchor_xyz)\n    Xt = torch.tensor(X, device=DEVICE)\n    E = float(energy_torch(Xt, segments, contacts=contacts).item())\n    bestX, bestE, bestw = X.copy(), E, w.copy()\n\n    T0, T1 = 0.04, 0.01\n\n    for it in range(steps):\n        temp = T0 + (T1 - T0) * (it / max(1, steps-1))\n\n        w_new = w.copy()\n        a, b = rng.integers(0, M), rng.integers(0, M)\n        if a == b:\n            b = (b + 1) % M\n        delta = float(rng.normal(0.0, 0.06))\n        w_new[a] = max(0.0, w_new[a] + delta)\n        w_new[b] = max(0.0, w_new[b] - delta)\n        ssum = w_new.sum()\n        if ssum < 1e-6:\n            w_new = w.copy()\n        else:\n            w_new = (w_new / ssum).astype(np.float32)\n\n        X_new = adaptive_rna_constraints(blend(w_new), segments, confidence=conf, passes=1,\n                                        anchor_mask=anchor_mask, anchor_xyz=anchor_xyz)\n        Xt_new = torch.tensor(X_new, device=DEVICE)\n        E_new = float(energy_torch(Xt_new, segments, contacts=contacts).item())\n\n        accept = (E_new <= E) or (rng.random() < np.exp(-(E_new - E) / max(1e-6, temp)))\n        if accept:\n            w, X, E = w_new, X_new, E_new\n            if E < bestE:\n                bestE, bestX, bestw = E, X.copy(), w.copy()\n\n    return bestX, bestE\n\n\n\ndef consensus_contacts(aligned_templates, sims, top_k=8, cutoff=14.0, min_sep=4, max_pairs=900):\n    # build voted contacts from top_k templates (weighted by similarity)\n    top_k = min(top_k, len(aligned_templates))\n    vote = defaultdict(float)\n    dist_sum = defaultdict(float)\n\n    for t in range(top_k):\n        X = aligned_templates[t]\n        w = max(0.0, float(sims[t]))  # weight by similarity\n\n        pairs = build_contacts_from_template(X, cutoff=cutoff, min_sep=min_sep, max_pairs=max_pairs)\n        for i, j, d0 in pairs:\n            key = (i, j)\n            vote[key] += w\n            dist_sum[key] += w * d0\n\n    # keep strongest voted contacts\n    items = [(k[0], k[1], dist_sum[k]/(vote[k]+1e-9), vote[k]) for k in vote.keys()]\n    items.sort(key=lambda x: x[3], reverse=True)\n\n    # return list of (i,j,d0) only\n    return [(i, j, d0) for (i, j, d0, _) in items[:max_pairs]]\n\n\n\n# -----------------------------\n# Build aligned templates for a target (expands K conformers, K usually 1)\n# -----------------------------\ndef build_aligned_templates(seq, cands, segments, M_targets=5):\n    picked = cands[:M_targets]\n    if not picked:\n        return [], [], None, None, None\n\n    adapted = []\n    sims = []\n    anchors = []\n\n    for (tid, tseq, sim, tcoordsK) in picked:  # tcoordsK: (K,Lt,3)\n        for m in range(tcoordsK.shape[0]):\n            X, a = adapt_template_to_query_dirfill(seq, tseq, tcoordsK[m], segments)\n            adapted.append(X)\n            anchors.append(a)\n            sims.append(sim)\n\n    if not adapted:\n        return [], [], None, None, None\n\n    ref = torch.tensor(adapted[0], device=DEVICE)\n    aligned = [adapted[0]]\n    for k in range(1, len(adapted)):\n        Q = torch.tensor(adapted[k], device=DEVICE)\n        Q_aligned = kabsch_align(ref, Q)\n        aligned.append(Q_aligned.detach().cpu().numpy().astype(np.float32))\n\n    anchor_mask = anchors[0]\n    anchor_xyz  = aligned[0].copy()\n   # contacts = build_contacts_from_template(aligned[0], cutoff=14.0, min_sep=4, max_pairs=700)\n    contacts = consensus_contacts(aligned, sims, top_k=8, cutoff=14.0, min_sep=4, max_pairs=900)\n\n    return aligned, sims, anchor_mask, anchor_xyz, contacts\n\n# -----------------------------\n# Predict 5 structures per target\n# -----------------------------\ndef predict_5_structures(row, train_seqs_df, train_coordsK_dict, top_n=60, M_targets=8, mc_steps=100):\n    tid = row[\"target_id\"]\n    seq = row[\"sequence\"]\n    assert set(seq).issubset(set(\"ACGU\")), f\"Non-ACGU in {tid}\"\n\n    segments = test_segs_map.get(tid, [(0, len(seq))])\n\n    cands = find_similar_sequences(seq, train_seqs_df, train_coordsK_dict, top_n=top_n)\n\n    if not cands:\n        base = np.zeros((len(seq), 3), dtype=np.float32)\n        for (s, e) in segments:\n            for i in range(s+1, e):\n                base[i] = base[i-1] + np.array([5.95, 0, 0], dtype=np.float32)\n        return [base.copy() for _ in range(5)]\n\n    aligned_templates, sims, anchor_mask, anchor_xyz, contacts = build_aligned_templates(\n        seq, cands, segments, M_targets=M_targets\n    )\n\n    if not aligned_templates:\n        t_id, t_seq, sim, tcoordsK = cands[0]\n        X0, a0 = adapt_template_to_query_dirfill(seq, t_seq, tcoordsK[0], segments)\n        X0 = adaptive_rna_constraints(X0, segments, confidence=sim, passes=3, anchor_mask=a0, anchor_xyz=X0.copy())\n        return [X0.copy() for _ in range(5)]\n\n    sim0 = float(max(sims)) if sims else 0.0\n    passes_out = 2 if sim0 >= 0.55 else 3\n\n    preds = []\n    energies = []\n\n    # generate 8 candidates, pick best 5 (best-of-5 friendly)\n    for i in range(8):\n        seed = (abs(hash(tid)) + 10007 * i) % (2**32)\n        rng = np.random.default_rng(seed)\n\n        # diversify by different contact subsets + slight noise\n        use_contacts = contacts\n        if contacts and i > 0:\n            n = len(contacts)\n            k = int(0.70 * n)\n            k = max(10, k)      # keep at least 10 if possible\n            k = min(n, k)       # never exceed population\n            if k < n:\n                keep = rng.choice(n, size=k, replace=False)\n                use_contacts = [contacts[t] for t in keep]\n            else:\n                use_contacts = contacts\n\n\n        noise_sigma = 0.0 if i == 0 else (0.10 + 0.03 * (i-1))  # small coordinate jitter\n\n        Xbest, Ebest = mc_blend_search(\n            adapted_list=aligned_templates,\n            sims=sims,\n            segments=segments,\n            seed=seed,\n            steps=mc_steps,\n            start=\"softmax\" if i < 4 else \"uniform\",\n            anchor_mask=anchor_mask,\n            anchor_xyz=anchor_xyz,\n            contacts=use_contacts,\n            noise_sigma=noise_sigma\n        )\n\n        Xbest = adaptive_rna_constraints(\n            Xbest, segments, confidence=sim0, passes=passes_out,\n            anchor_mask=anchor_mask, anchor_xyz=anchor_xyz\n        )\n\n        preds.append(Xbest)\n        energies.append(Ebest)\n\n    order = np.argsort(energies)\n    chosen = []\n\n    def rmsd(a, b):\n        return float(np.sqrt(np.mean((a - b) ** 2)))\n\n    for idx in order:\n        X = preds[idx]\n        if all(rmsd(X, Y) >= 0.35 for Y in chosen):\n            chosen.append(X)\n        if len(chosen) == 5:\n            break\n\n    if len(chosen) < 5:\n        for idx in order:\n            if len(chosen) == 5:\n                break\n            X = preds[idx]\n            if all(not np.allclose(X, Y) for Y in chosen):\n                chosen.append(X)\n\n    return chosen[:5]\n\n# -----------------------------\n# Generate submission\n# -----------------------------\nall_rows = []\nt0 = time.time()\n\nfor idx, row in test_seqs.iterrows():\n    if idx % 5 == 0:\n        print(f\"[{idx}/{len(test_seqs)}] elapsed {time.time()-t0:.1f}s | device={DEVICE}\")\n\n    preds = predict_5_structures(\n        row,\n        train_seqs_df=train_seqs,\n        train_coordsK_dict=train_coordsK_dict,\n        top_n=30,\n        M_targets=5,\n        mc_steps=80,\n    )\n\n    tid = row[\"target_id\"]\n    seq = row[\"sequence\"]\n\n    for j in range(len(seq)):\n        r = {\"ID\": f\"{tid}_{j+1}\", \"resname\": seq[j], \"resid\": j+1}\n        for i in range(5):\n            r[f\"x_{i+1}\"], r[f\"y_{i+1}\"], r[f\"z_{i+1}\"] = preds[i][j].tolist()\n        all_rows.append(r)\n\nsub = pd.DataFrame(all_rows)\ncols = [\"ID\", \"resname\", \"resid\"] + [f\"{c}_{i}\" for i in range(1, 6) for c in [\"x\", \"y\", \"z\"]]\n\ncoord_cols = [c for c in cols if c.startswith((\"x_\", \"y_\", \"z_\"))]\nsub[coord_cols] = sub[coord_cols].clip(-999.999, 9999.999)\n\nsub[cols].to_csv(\"submission.csv\", index=False)\nprint(\"Saved submission.csv\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-12T12:41:43.650903Z","iopub.execute_input":"2026-02-12T12:41:43.651278Z","iopub.status.idle":"2026-02-12T12:46:15.729926Z","shell.execute_reply.started":"2026-02-12T12:41:43.651229Z","shell.execute_reply":"2026-02-12T12:46:15.729134Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}