{"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":"none","dataSources":[{"sourceId":118765,"databundleVersionId":15231210,"sourceType":"competition"},{"sourceId":14441699,"sourceType":"datasetVersion","datasetId":9224635},{"sourceId":14604295,"sourceType":"datasetVersion","datasetId":9328538}],"dockerImageVersionId":31259,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false},"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/datasets/kami1976/biopython-cp312/biopython-1.86-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T14:34:40.748137Z","iopub.execute_input":"2026-02-10T14:34:40.749452Z","iopub.status.idle":"2026-02-10T14:34:47.859655Z","shell.execute_reply.started":"2026-02-10T14:34:40.749402Z","shell.execute_reply":"2026-02-10T14:34:47.858443Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Stanford RNA 3D Folding Part 2 — Deterministic Template Baseline\n# CPU-only, no external data, no ML training\n# Output: submission.csv (5 coordinate sets per residue)\n# ============================================================\n\nimport pandas as pd\nimport numpy as np\nimport random\nimport time\nimport warnings\nimport os, sys\n\nos.environ.setdefault(\"OMP_NUM_THREADS\", \"1\")\nos.environ.setdefault(\"MKL_NUM_THREADS\", \"1\")\nos.environ.setdefault(\"OPENBLAS_NUM_THREADS\", \"1\")\nos.environ.setdefault(\"NUMEXPR_NUM_THREADS\", \"1\")\n\nrandom.seed(0)\nnp.random.seed(0)\n\n# Determinism hygiene\nos.environ[\"PYTHONHASHSEED\"] = \"0\"\nwarnings.filterwarnings(\"ignore\")\n\nimport zlib\n\nVALID_NTS = set(\"ACGU\")\n\ndef stable_u32(s: str, salt: int = 0) -> int:\n    \"\"\"\n    Stable, cross-run deterministic 32-bit seed.\n    Avoids Python's hash() nondeterminism (even if PYTHONHASHSEED is set too late).\n    \"\"\"\n    b = (str(s) + \"|\" + str(int(salt))).encode(\"utf-8\")\n    return zlib.crc32(b) & 0xFFFFFFFF\n\ndef assert_acgu_only(seq: str, ctx: str = \"\"):\n    seq = str(seq)\n    bad = set(seq) - VALID_NTS\n    if bad:\n        raise ValueError(f\"[ACGU-VALIDATION-FAIL] {ctx} bad={sorted(bad)}\")\n\n# -----------------------------\n# A) Data ingestion\n# -----------------------------\nDATA_PATH = \"/kaggle/input/stanford-rna-3d-folding-2/\"\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\nsys.path.append(os.path.join(DATA_PATH, \"extra\"))\n\n# --- Robust FASTA parser: prefer dataset extra/parse_fasta_py.py; fallback if import fails ---\ntry:\n    import typing as _typing\n    import builtins as _builtins\n\n    _builtins.Dict  = getattr(_typing, \"Dict\")\n    _builtins.Tuple = getattr(_typing, \"Tuple\")\n    _builtins.List  = getattr(_typing, \"List\")\n\n    from parse_fasta_py import parse_fasta as _parse_fasta_raw\n\n    # Normalize output to: {chain_id: sequence_string}\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    # Fallback FASTA parser: {chain_id: sequence_string}\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]  # first token as chain id\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    \"\"\"\n    Returns list of (start,end) segments in row['sequence'] corresponding to chain copies\n    in stoichiometry order. If parsing fails or mismatch: fallback single segment (0,L).\n    \"\"\"\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)  # chain_id -> sequence\n        order = parse_stoichiometry(stoich)\n\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\n        if pos != len(seq):\n            return [(0, len(seq))]\n        return segs\n    except Exception:\n        return [(0, len(seq))]\n\ndef build_segments_map(df):\n    seg_map = {}\n    stoich_map = {}\n    for _, r in df.iterrows():\n        tid = r[\"target_id\"]\n        seg_map[tid] = get_chain_segments(r)\n        stoich_map[tid] = str(r.get(\"stoichiometry\", \"\") if not pd.isna(r.get(\"stoichiometry\", \"\")) else \"\")\n    return seg_map, stoich_map\n\ndef get_chain_list(row):\n    \"\"\"\n    Returns:\n      chain_seqs: [seq_chain_copy_0, seq_chain_copy_1, ...] in stoichiometry order\n      chain_ids:  [id_copy_0, id_copy_1, ...] aligned with chain_seqs (ids include copy index)\n    Fallback: single chain [row['sequence']]\n    \"\"\"\n    seq = str(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 [seq], [\"CHAIN0\"]\n\n    try:\n        chain_dict = parse_fasta(all_seq)  # chain_id -> sequence\n        order = parse_stoichiometry(stoich)\n\n        chain_seqs = []\n        chain_ids = []\n        for ch, cnt in order:\n            base = chain_dict.get(ch)\n            if base is None:\n                return [seq], [\"CHAIN0\"]\n            for k in range(cnt):\n                chain_seqs.append(str(base))\n                chain_ids.append(f\"{ch}_{k+1}\")\n        # sanity: concatenation length must match\n        if sum(len(x) for x in chain_seqs) != len(seq):\n            return [seq], [\"CHAIN0\"]\n        return chain_seqs, chain_ids\n    except Exception:\n        return [seq], [\"CHAIN0\"]\n\ndef build_chain_list_map(df):\n    chain_list_map = {}\n    chain_id_map = {}\n    for _, r in df.iterrows():\n        tid = r[\"target_id\"]\n        chain_seqs, chain_ids = get_chain_list(r)\n        chain_list_map[tid] = chain_seqs\n        chain_id_map[tid] = chain_ids\n    return chain_list_map, chain_id_map\n\ntrain_chain_list_map, train_chain_ids_map = build_chain_list_map(train_seqs)\ntest_chain_list_map,  test_chain_ids_map  = build_chain_list_map(test_seqs)\n\ntrain_segs_map, train_stoich_map = build_segments_map(train_seqs)\ntest_segs_map,  test_stoich_map  = build_segments_map(test_seqs)\n\n# -----------------------------\n# B) Labels to templates\n# -----------------------------\ndef process_labels(labels_df: pd.DataFrame):\n    \"\"\"\n    train_coords_dict: {target_id: (L,3) coords from x_1,y_1,z_1}\n    Group key = labels_df['ID'] split at last underscore\n    \"\"\"\n    coords_dict = {}\n    prefixes = labels_df[\"ID\"].str.rsplit(\"_\", n=1).str[0]\n    for id_prefix, group in labels_df.groupby(prefixes, sort=False):\n        coords_dict[id_prefix] = group.sort_values(\"resid\")[[\"x_1\", \"y_1\", \"z_1\"]].values\n    return coords_dict\n\ntrain_coords_dict = process_labels(train_labels)\n\n# -----------------------------\n# C) Similarity search (PairwiseAligner strict global w/ strong gaps)\n# -----------------------------\nfrom Bio.Align import PairwiseAligner\n\naligner = PairwiseAligner()\naligner.mode = \"global\"\naligner.match_score = 2\naligner.mismatch_score = -1.5\naligner.open_gap_score = -8\naligner.extend_gap_score = -0.4\n\n# Explicit terminal gap penalties (avoid end-gap sliding)\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\nimport itertools\n\ndef _norm_align_score(a: str, b: str) -> float:\n    # same normalization logic you used, but per-chain\n    raw = aligner.score(a, b)\n    denom = 2.0 * max(1, min(len(a), len(b)))\n    return float(raw) / denom\n\n\ndef chain_assignment_similarity(q_chains, t_chains):\n    \"\"\"\n    Returns:\n      sim: float in roughly [-1,1] (normalized)\n      mapping: list of (qi, tj) chosen to maximize total score\n    Deterministic: brute force for <=6 chains; greedy fallback otherwise.\n    \"\"\"\n    nq, nt = len(q_chains), len(t_chains)\n    if nq == 0 or nt == 0:\n        return -1.0, []\n\n    # matrix of per-chain normalized scores\n    S = np.zeros((nq, nt), dtype=float)\n    for i in range(nq):\n        for j in range(nt):\n            S[i, j] = _norm_align_score(q_chains[i], t_chains[j])\n\n    m = min(nq, nt)\n\n    # brute force assignment if small\n    if max(nq, nt) <= 6:\n        best = (-1e18, None)\n        # choose m template indices and permute\n        for t_idx_subset in itertools.combinations(range(nt), m):\n            for perm in itertools.permutations(t_idx_subset, m):\n                tot = 0.0\n                for qi in range(m):\n                    tot += S[qi, perm[qi]]\n                if tot > best[0]:\n                    best = (tot, perm)\n        perm = best[1]\n        mapping = [(qi, int(perm[qi])) for qi in range(m)]\n        sim = best[0] / max(1, m)\n        return float(sim), mapping\n\n    # greedy fallback: repeatedly pick best remaining pair\n    mapping = []\n    used_q = set()\n    used_t = set()\n    for _ in range(m):\n        best_val = -1e18\n        best_pair = None\n        for i in range(nq):\n            if i in used_q:\n                continue\n            for j in range(nt):\n                if j in used_t:\n                    continue\n                v = S[i, j]\n                if v > best_val:\n                    best_val = v\n                    best_pair = (i, j)\n        if best_pair is None:\n            break\n        used_q.add(best_pair[0])\n        used_t.add(best_pair[1])\n        mapping.append(best_pair)\n\n    sim = np.mean([S[i, j] for (i, j) in mapping]) if mapping else -1.0\n    return float(sim), mapping\n\n\ndef find_similar_sequences(\n    query_seq: str,\n    query_chain_list: list,\n    train_seqs_df: pd.DataFrame,\n    train_coords_dict: dict,\n    top_n: int = 5,\n):\n    \"\"\"\n    Returns list of tuples:\n      (target_id, train_seq, sim, train_coords, train_segments, train_chain_list, chain_map)\n    where chain_map is list of (q_chain_index, t_chain_index)\n    \"\"\"\n    similar = []\n\n    for _, row in train_seqs_df.iterrows():\n        tid = row[\"target_id\"]\n        t_seq = row[\"sequence\"]\n        if tid not in train_coords_dict:\n            continue\n\n        # cheap length filter on total length\n        if abs(len(t_seq) - len(query_seq)) / max(len(t_seq), len(query_seq)) > 0.3:\n            continue\n\n        t_chain_list = train_chain_list_map.get(tid, [t_seq])\n        t_segments = train_segs_map.get(tid, [(0, len(t_seq))])\n\n        # prefer matching #chains; still allow mismatch but penalize\n        sim, cmap = chain_assignment_similarity(query_chain_list, t_chain_list)\n        chain_penalty = 0.0\n        if len(query_chain_list) != len(t_chain_list):\n            chain_penalty = 0.12 * abs(len(query_chain_list) - len(t_chain_list))\n        sim = sim - chain_penalty\n\n        similar.append((tid, t_seq, sim, train_coords_dict[tid], t_segments, t_chain_list, cmap))\n\n    similar.sort(key=lambda x: x[2], reverse=True)\n    return similar[:top_n]\n\n# -----------------------------\n# D) Template transfer (alignment.aligned block mapping + interpolation fill)\n# -----------------------------\ndef adapt_template_to_query(query_seq: str, template_seq: str, template_coords: np.ndarray):\n    alignment = next(iter(aligner.align(query_seq, template_seq)))\n    new_coords = np.full((len(query_seq), 3), np.nan, dtype=float)\n\n    # Vectorized chunk mapping via aligned blocks\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    # Fill unmatched residues by interpolation / edge-fill / fallback line\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\n            if prev_v >= 0 and next_v >= 0:\n                w = (i - prev_v) / (next_v - prev_v)\n                new_coords[i] = (1 - w) * new_coords[prev_v] + w * new_coords[next_v]\n            elif prev_v >= 0:\n                new_coords[i] = new_coords[prev_v] + [5.95, 0, 0]\n            elif next_v >= 0:\n                new_coords[i] = new_coords[next_v] + [5.95, 0, 0]\n            else:\n                new_coords[i] = [i * 5.95, 0, 0]\n\n    return np.nan_to_num(new_coords)\n\ndef _adapt_chain_coords(q_chain_seq: str, t_chain_seq: str, t_chain_coords: np.ndarray):\n    \"\"\"\n    Adapt template coords for one chain (alignment blocks + within-chain interpolation fill).\n    Returns (Lq,3).\n    \"\"\"\n    alignment = next(iter(aligner.align(q_chain_seq, t_chain_seq)))\n    out = np.full((len(q_chain_seq), 3), np.nan, dtype=float)\n\n    for (qs, qe), (ts, te) in zip(*alignment.aligned):\n        chunk = t_chain_coords[ts:te]\n        if len(chunk) == (qe - qs):\n            out[qs:qe] = chunk\n\n    # fill within-chain only\n    for i in range(len(out)):\n        if np.isnan(out[i, 0]):\n            prev_v = next((j for j in range(i - 1, -1, -1) if not np.isnan(out[j, 0])), -1)\n            next_v = next((j for j in range(i + 1, len(out)) if not np.isnan(out[j, 0])), -1)\n\n            if prev_v >= 0 and next_v >= 0:\n                w = (i - prev_v) / (next_v - prev_v)\n                out[i] = (1 - w) * out[prev_v] + w * out[next_v]\n            elif prev_v >= 0:\n                out[i] = out[prev_v] + [5.95, 0, 0]\n            elif next_v >= 0:\n                out[i] = out[next_v] + [5.95, 0, 0]\n            else:\n                out[i] = [i * 5.95, 0, 0]\n\n    return np.nan_to_num(out)\n\n\ndef adapt_template_to_query_multichain(\n    query_seq: str,\n    query_segments: list,\n    query_chain_list: list,\n    template_seq: str,\n    template_segments: list,\n    template_chain_list: list,\n    template_coords: np.ndarray,\n    chain_map: list,\n):\n    \"\"\"\n    Build a full (Lq,3) coordinate array by adapting per chain.\n    Uses chain_map: list of (q_idx, t_idx).\n    Any query chain without a mapped template chain -> straight line initialized near origin\n    (then later docking jitter/refinement will separate).\n    \"\"\"\n    Lq = len(query_seq)\n    out = np.zeros((Lq, 3), dtype=float)\n\n    # default: straight line per query chain (so unmapped chains are valid)\n    for qi, (qs, qe) in enumerate(query_segments):\n        for j in range(qs + 1, qe):\n            out[j] = out[j - 1] + [5.95, 0, 0]\n\n    # apply mapped chains\n    for (qi, ti) in chain_map:\n        if qi < 0 or qi >= len(query_segments):\n            continue\n        if ti < 0 or ti >= len(template_segments):\n            continue\n\n        qs, qe = query_segments[qi]\n        ts, te = template_segments[ti]\n\n        q_chain_seq = query_chain_list[qi]\n        t_chain_seq = template_chain_list[ti]\n\n        t_chain_coords = template_coords[ts:te]\n        # safety: ensure template slice length matches template chain seq\n        if len(t_chain_coords) != len(t_chain_seq):\n            # fallback: avoid crash, keep default chain line\n            continue\n\n        adapted_chain = _adapt_chain_coords(q_chain_seq, t_chain_seq, t_chain_coords)\n\n        # write into global frame at the correct segment\n        if (qe - qs) == len(adapted_chain):\n            out[qs:qe] = adapted_chain\n\n    return out\n\n\n\n# -----------------------------\n# E) Segment-aware local refinement (US-align compatible)\n# -----------------------------\ndef adaptive_rna_constraints(coordinates: np.ndarray, target_id: str, confidence: float = 1.0, passes: int = 2):\n    \"\"\"\n    Apply within each chain segment only (no smoothing across chain breaks).\n    - i,i+1 bond target ~5.95 Å (symmetric)\n    - i,i+2 target ~10.2 Å (symmetric)\n    - small Laplacian smoothing\n    - light self-avoidance on subsampled points\n    Strength increases when confidence is lower:\n      strength = max(0.02, 0.75*(1-min(conf,0.90)))\n    \"\"\"\n    coords = coordinates.copy()\n    segments = test_segs_map.get(target_id, [(0, len(coords))])\n\n    strength = 0.75 * (1.0 - min(confidence, 0.90))\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            # (1) bond i,i+1 to ~5.95Å\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.22 * strength)\n            X[:-1] -= adj\n            X[1:]  += adj\n\n            # (2) soft i,i+2 to ~10.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.10 * strength)\n            X[:-2] -= adj2\n            X[2:]  += adj2\n\n            # (3) Laplacian smoothing\n            lap = 0.5 * (X[:-2] + X[2:]) - X[1:-1]\n            X[1:-1] += (0.06 * strength) * lap\n\n            # (4) light self-avoidance (subsample)\n            if L >= 25:\n                k = min(L, 160) if L > 220 else L\n                idx = np.linspace(0, L - 1, k).astype(int) if k < L else np.arange(L)\n\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.2)\n                if np.any(mask):\n                    force = (3.2 - distm) / distm\n                    vec = (diff * force[:, :, None] * mask[:, :, None]).sum(axis=1)\n                    X[idx] += (0.015 * strength) * vec\n\n            coords[s:e] = X\n\n    return coords\n\ndef relax_interchain_clashes(coords: np.ndarray, segments: list, strength: float = 0.25, iters: int = 2):\n    \"\"\"\n    Mild repulsion between chains if their subsampled points overlap too much.\n    Deterministic: no RNG.\n    \"\"\"\n    X = coords.copy()\n    if len(segments) <= 1:\n        return X\n\n    for _ in range(iters):\n        centers = []\n        subs = []\n        for (s, e) in segments:\n            P = X[s:e]\n            c = P.mean(axis=0)\n            centers.append(c)\n            L = e - s\n            k = min(L, 80)\n            idx = np.linspace(0, L - 1, k).astype(int) if k < L else np.arange(L)\n            subs.append(P[idx])\n\n        centers = np.array(centers, float)\n\n        # pairwise chain repulsion if too close\n        for a in range(len(segments)):\n            for b in range(a + 1, len(segments)):\n                Pa, Pb = subs[a], subs[b]\n                diff = Pa[:, None, :] - Pb[None, :, :]\n                dist = np.linalg.norm(diff, axis=2)\n                dmin = float(np.min(dist)) if dist.size else 1e9\n\n                if dmin < 2.4:\n                    va = centers[a] - centers[b]\n                    n = np.linalg.norm(va) + 1e-12\n                    dirv = va / n\n                    push = (2.4 - dmin) * strength\n\n                    sa, ea = segments[a]\n                    sb, eb = segments[b]\n                    X[sa:ea] += dirv * (0.5 * push)\n                    X[sb:eb] -= dirv * (0.5 * push)\n\n    return X\n\n# -----------------------------\n# F) Best-of-5 predictions (deterministic seed)\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        [\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        ],\n        dtype=float,\n    )\n\ndef apply_hinge(coords, seg, rng, max_angle_deg=25):\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\ndef jitter_chains(coords, segments, rng, max_angle_deg=12, max_trans=1.5):\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\ndef smooth_wiggle(coords, segments, rng, amp=0.8):\n    X = coords.copy()\n    for (s, e) in segments:\n        L = e - s\n        if L < 20:\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\ndef predict_rna_structures(row, train_seqs_df, train_coords_dict, n_predictions=5):\n    tid = row[\"target_id\"]\n    seq = row[\"sequence\"]\n\n    # Canonical A/C/G/U only (do not remap)\n    assert_acgu_only(seq, ctx=tid)\n    segments = test_segs_map.get(tid, [(0, len(seq))])\n    env = extract_env_context(row)\n    ctx = context_strength(env, tid)  \n\n    # Candidate pool top_n=30 then sample for diversity\n    q_chain_list = test_chain_list_map.get(tid, [seq])\n    q_segments = test_segs_map.get(tid, [(0, len(seq))])\n\n    cands = find_similar_sequences(\n        query_seq=seq,\n        query_chain_list=q_chain_list,\n        train_seqs_df=train_seqs_df,\n        train_coords_dict=train_coords_dict,\n        top_n=30,\n    )\n\n    # tuple: (t_id, t_seq, sim, t_coords, t_segments, t_chain_list, chain_map)\n    assert all(len(c[3]) == len(c[1]) for c in cands), \"Template coords/seq length mismatch\"\n\n    \n    predictions = []\n    used = set()\n\n    for i in range(n_predictions):\n        seed = stable_u32(tid, salt=i * 10007)\n        rng = np.random.default_rng(seed)\n\n        if not cands:\n            # Hard fallback: straight line per chain segment\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        # Template choice\n        if i == 0:\n            t_id, t_seq, sim, t_coords, t_segments, t_chain_list, chain_map = cands[0]\n        else:\n            K = min(12, len(cands))\n            sims = np.array([cands[k][2] for k in range(K)], float)\n            w = np.exp((sims - sims.max()) / 0.08)\n            for k in range(K):\n                if cands[k][0] in used:\n                    w[k] *= 0.10\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, t_segments, t_chain_list, chain_map = cands[k]\n\n        used.add(t_id)\n\n        if len(q_chain_list) == 1 and len(t_chain_list) == 1:\n            adapted = adapt_template_to_query(query_seq=seq, template_seq=t_seq, template_coords=t_coords)\n        else:\n            adapted = adapt_template_to_query_multichain(\n                query_seq=seq,\n                query_segments=q_segments,\n                query_chain_list=q_chain_list,\n                template_seq=t_seq,\n                template_segments=t_segments,\n                template_chain_list=t_chain_list,\n                template_coords=t_coords,\n                chain_map=chain_map,\n        )\n\n        # Context increases diversity amplitude (multi-state coverage)\n        amp = 1.0 + 1.25 * ctx  # ctx==0 -> 1.0, ctx==1 -> 2.25\n\n        if i == 0:\n            X = adapted\n        elif i == 1:\n            X = adapted + rng.normal(0, amp * max(0.01, (0.40 - sim) * 0.06), adapted.shape)\n        elif i == 2:\n            # state kick (localized) + optional hinge\n            X = local_state_kick(adapted, segments, rng, amp=min(2.0, amp))\n            longest = max(segments, key=lambda se: se[1] - se[0])\n            X = apply_hinge(X, longest, rng, max_angle_deg=22 * amp)\n        elif i == 3:\n            # alternative docking mode: rotate/translate chains more when context exists\n            X = jitter_chains(adapted, segments, rng, max_angle_deg=10 * amp, max_trans=1.0 * amp)\n        else:\n            X = smooth_wiggle(adapted, segments, rng, amp=0.7 * amp)\n\n        refined = adaptive_rna_constraints(X, tid, confidence=sim, passes=2)\n        refined = relax_interchain_clashes(refined, segments, strength=0.22, iters=2)\n        predictions.append(refined)\n\n    return predictions\n\ndef extract_env_context(row: pd.Series):\n    \"\"\"\n    Robustly extract ligand/protein context if columns exist.\n    Returns dict with strings (possibly empty).\n    \"\"\"\n    # common candidate column names (safe if absent)\n    ligand_cols = [\"ligand_smiles\", \"smiles\", \"SMILES\", \"ligand\", \"ligand_context\"]\n    prot_cols   = [\"protein_fasta\", \"protein_sequence\", \"protein_seq\", \"partner_protein_fasta\", \"partner_protein\"]\n\n    lig = \"\"\n    prot = \"\"\n\n    for c in ligand_cols:\n        if c in row.index and not pd.isna(row[c]):\n            lig = str(row[c]).strip()\n            break\n    for c in prot_cols:\n        if c in row.index and not pd.isna(row[c]):\n            prot = str(row[c]).strip()\n            break\n\n    return {\"ligand\": lig, \"protein\": prot}\n\n\ndef context_strength(env: dict, tid: str):\n    \"\"\"\n    Deterministic scalar in [0,1] indicating how strongly to diversify.\n    \"\"\"\n    key = (env.get(\"ligand\", \"\") + \"||\" + env.get(\"protein\", \"\")).strip()\n    if key == \"\":\n        return 0.0\n    # stable pseudo-random but deterministic\n    u = stable_u32(tid + \"|\" + key, salt=777)\n    return float((u % 1000) / 1000.0)\n\n\ndef local_state_kick(coords: np.ndarray, segments: list, rng: np.random.Generator, amp: float = 1.0):\n    \"\"\"\n    Create a deterministic, localized conformational shift inside one chosen chain segment.\n    \"\"\"\n    X = coords.copy()\n    if not segments:\n        return X\n\n    # choose one segment and one pivot window deterministically via rng\n    seg = segments[int(rng.integers(0, len(segments)))]\n    s, e = seg\n    L = e - s\n    if L < 25:\n        return X\n\n    pivot = s + int(rng.integers(8, L - 8))\n    axis = rng.normal(size=3)\n    ang = np.deg2rad(float(rng.uniform(-18, 18))) * amp\n    R = _rotmat(axis, ang)\n\n    p0 = X[pivot].copy()\n    X[pivot:e] = (X[pivot:e] - p0) @ R.T + p0\n\n    # add a smooth wiggle to that chain only\n    n_ctrl = 5\n    ctrl_x = np.linspace(0, L - 1, n_ctrl)\n    ctrl_disp = rng.normal(0, 0.55 * 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\n# -----------------------------\n# G) Submission writer (exact schema)\n# -----------------------------\nall_predictions = []\nstart_time = time.time()\n\nfor idx, row in test_seqs.iterrows():\n    if idx % 10 == 0:\n        print(f\"Processing {idx} | {time.time() - start_time:.1f}s\")\n    tid = row[\"target_id\"]\n    seq = row[\"sequence\"]\n\n    preds = predict_rna_structures(row, train_seqs, train_coords_dict, n_predictions=5)\n\n    # Safety: each prediction must be (L,3)\n    L = len(seq)\n    for p in preds:\n        assert isinstance(p, np.ndarray) and p.shape == (L, 3), f\"Bad pred shape for {tid}: {getattr(p,'shape',None)}\"\n        assert np.isfinite(p).all(), f\"Non-finite coords in {tid}\"\n\n    for j in range(L):\n        res = {\"ID\": f\"{tid}_{j+1}\", \"resname\": seq[j], \"resid\": j + 1}\n        for i in range(5):\n            res[f\"x_{i+1}\"], res[f\"y_{i+1}\"], res[f\"z_{i+1}\"] = preds[i][j]\n        all_predictions.append(res)\n\nsub = pd.DataFrame(all_predictions)\n\ncols = [\"ID\", \"resname\", \"resid\"] + [f\"{c}_{i}\" for i in range(1, 6) for c in [\"x\", \"y\", \"z\"]]\n\n# Clip explicitly (competition clips coords; prevent explosions)\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(\"submission.csv! saved\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T14:34:47.86231Z","iopub.execute_input":"2026-02-10T14:34:47.863006Z","iopub.status.idle":"2026-02-10T14:37:43.500697Z","shell.execute_reply.started":"2026-02-10T14:34:47.862968Z","shell.execute_reply":"2026-02-10T14:37:43.499685Z"}},"outputs":[],"execution_count":null}]}