{"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":"gpu","dataSources":[{"sourceId":118765,"databundleVersionId":15231210,"sourceType":"competition"},{"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-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-05T15:23:16.089246Z","iopub.execute_input":"2026-02-05T15:23:16.089637Z","iopub.status.idle":"2026-02-05T15:23:21.264487Z","shell.execute_reply.started":"2026-02-05T15:23:16.089611Z","shell.execute_reply":"2026-02-05T15:23:21.263557Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Template-based RNA Structure Prediction (FULL UPDATED CODE)\n# - Deterministic seeds (stable_hash32)\n# - Fast template retrieval (train library + k-mer Jaccard prefilter)\n# - Best-of-5 diverse predictions + chain-aware constraints\n# - Faster submission build (concat dfs, no per-residue dict loop)\n# ============================================================\n\nimport pandas as pd\nimport numpy as np\nimport random\nimport time\nimport warnings\nimport os, sys\nimport zlib\n\nwarnings.filterwarnings(\"ignore\")\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# -----------------------------\n# Deterministic hashing\n# -----------------------------\ndef stable_hash32(s: str) -> int:\n    return zlib.adler32(s.encode(\"utf-8\")) & 0xFFFFFFFF\n\n# -----------------------------\n# Robust FASTA parser\n# -----------------------------\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    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    \"\"\"\n    Returns list of (start,end) segments in row['sequence'] corresponding to chain copies in stoichiometry order.\n    Falls back to single segment if parsing fails.\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)\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    stoich_map = {}\n    for r in df.itertuples(index=False):\n        tid = r.target_id\n        # itertuples gives attributes; we still need row dict for get_chain_segments\n        # so use df.loc? -> avoid slow: use df[df.target_id==] no.\n        # simplest: fallback to iterrows for this one-time build (cheap).\n        pass\n    # one-time build (cost acceptable)\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\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# Labels -> coords dict\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, 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# Aligner\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\n# ============================================================\n# FAST TEMPLATE RETRIEVAL (train library + k-mer prefilter)\n# ============================================================\nKMER_K = 5\nPREFILTER_TOP = 250   # survive cheap k-mer filter\nALIGN_TOP_N = 30      # survive aligner.score\n\n_BASE2 = {\"A\": 0, \"C\": 1, \"G\": 2, \"U\": 3}\n\ndef kmer_set_2bit(seq: str, k: int = KMER_K):\n    if len(seq) < k:\n        return frozenset()\n    mask = (1 << (2 * k)) - 1\n    code = 0\n    out = set()\n    for i in range(k):\n        code = ((code << 2) | _BASE2[seq[i]]) & mask\n    out.add(code)\n    for i in range(k, len(seq)):\n        code = ((code << 2) | _BASE2[seq[i]]) & mask\n        out.add(code)\n    return frozenset(out)\n\n# Build train library once (coords-aligned only)\nTRAIN_IDS, TRAIN_SEQS, TRAIN_COORDS, TRAIN_KMERS = [], [], [], []\n_tmp_lens = []\n\nfor r in train_seqs.itertuples(index=False):\n    tid = r.target_id\n    if tid not in train_coords_dict:\n        continue\n    seq = r.sequence\n    coords = train_coords_dict[tid]\n    if len(coords) != len(seq):\n        continue\n    TRAIN_IDS.append(tid)\n    TRAIN_SEQS.append(seq)\n    TRAIN_COORDS.append(coords)\n    _tmp_lens.append(len(seq))\n    TRAIN_KMERS.append(kmer_set_2bit(seq, KMER_K))\n\nTRAIN_LENS = np.asarray(_tmp_lens, dtype=np.int32)\nprint(f\"Train library: {len(TRAIN_IDS)} templates (coords-aligned)\")\n\ndef find_similar_sequences(query_seq: str, top_n: int = ALIGN_TOP_N, prefilter_top: int = PREFILTER_TOP):\n    qlen = len(query_seq)\n    qkm = kmer_set_2bit(query_seq, KMER_K)\n    qkm_len = len(qkm)\n\n    if len(TRAIN_IDS) == 0:\n        return []\n\n    L = TRAIN_LENS\n    maxL = np.maximum(L, qlen)\n    keep = (np.abs(L - qlen) / maxL) <= 0.30\n    idxs = np.where(keep)[0]\n    if idxs.size == 0:\n        return []\n\n    scored = []\n    for i in idxs:\n        tkm = TRAIN_KMERS[i]\n        if qkm_len == 0 or len(tkm) == 0:\n            jac = 0.0\n        else:\n            inter = len(qkm & tkm)\n            union = qkm_len + len(tkm) - inter\n            jac = inter / union if union else 0.0\n        scored.append((jac, i))\n\n    scored.sort(key=lambda x: x[0], reverse=True)\n    scored = scored[:min(prefilter_top, len(scored))]\n\n    sims = []\n    denom_base = 2.0 * qlen  # not exactly used; we compute denom per-template\n    for jac, i in scored:\n        t_seq = TRAIN_SEQS[i]\n        raw = aligner.score(query_seq, t_seq)\n        denom = (2.0 * min(qlen, len(t_seq)))\n        norm = float(raw / denom) if denom > 0 else -1.0\n        sims.append((TRAIN_IDS[i], t_seq, norm, TRAIN_COORDS[i]))\n\n    sims.sort(key=lambda x: x[2], reverse=True)\n    return sims[:top_n]\n\n# ============================================================\n# Template -> query coordinate adaptation\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    # interpolation/fallback\n    for i in range(len(new_coords)):\n        if np.isnan(new_coords[i, 0]):\n            prev_v = next((j for j in range(i - 1, -1, -1) if not np.isnan(new_coords[j, 0])), -1)\n            next_v = next((j for j in range(i + 1, len(new_coords)) if not np.isnan(new_coords[j, 0])), -1)\n            if prev_v >= 0 and next_v >= 0:\n                w = (i - prev_v) / (next_v - prev_v)\n                new_coords[i] = (1 - w) * new_coords[prev_v] + w * new_coords[next_v]\n            elif prev_v >= 0:\n                new_coords[i] = new_coords[prev_v] + [3, 0, 0]\n            elif next_v >= 0:\n                new_coords[i] = new_coords[next_v] + [3, 0, 0]\n            else:\n                new_coords[i] = [i * 3, 0, 0]\n\n    return np.nan_to_num(new_coords)\n\n# ============================================================\n# Chain-aware constraints + diversity transforms\n# ============================================================\ndef adaptive_rna_constraints(coordinates, target_id, confidence=1.0, passes=2):\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            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            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            lap = 0.5 * (X[:-2] + X[2:]) - X[1:-1]\n            X[1:-1] += (0.06 * strength) * lap\n\n            if L >= 25:\n                k = min(L, 160) if L > 220 else L\n                idx = np.linspace(0, L - 1, k).astype(int) if k < L else np.arange(L)\n\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 _rotmat(axis, ang):\n    axis = np.asarray(axis, float)\n    axis = axis / (np.linalg.norm(axis) + 1e-12)\n    x, y, z = axis\n    c, s = np.cos(ang), np.sin(ang)\n    C = 1.0 - c\n    return np.array([\n        [c + x*x*C,     x*y*C - z*s, x*z*C + y*s],\n        [y*x*C + z*s,   c + y*y*C,   y*z*C - x*s],\n        [z*x*C - y*s,   z*y*C + x*s, c + z*z*C]\n    ], dtype=float)\n\ndef apply_hinge(coords, seg, rng, max_angle_deg=25):\n    s, e = seg\n    L = e - s\n    if L < 30:\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\n# ============================================================\n# Predictor (best-of-5)\n# ============================================================\ndef predict_rna_structures(tid: str, seq: str, n_predictions: int = 5):\n    assert set(seq).issubset(set(\"ACGU\")), f\"Non-ACGU in {tid}; do not remap here.\"\n    segments = test_segs_map.get(tid, [(0, len(seq))])\n\n    cands = find_similar_sequences(query_seq=seq, top_n=ALIGN_TOP_N)\n    predictions = []\n    used = set()\n\n    for i in range(n_predictions):\n        seed = (stable_hash32(tid) + i * 10007) & 0xFFFFFFFF\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(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 = cands[k]\n\n        used.add(t_id)\n\n        adapted = adapt_template_to_query(query_seq=seq, template_seq=t_seq, template_coords=t_coords)\n\n        if i == 0:\n            X = adapted\n        elif i == 1:\n            X = adapted + rng.normal(0, max(0.01, (0.40 - sim) * 0.06), 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=22)\n        elif i == 3:\n            X = jitter_chains(adapted, segments, rng, max_angle_deg=10, max_trans=1.0)\n        else:\n            X = smooth_wiggle(adapted, segments, rng, amp=0.7)\n\n        refined = adaptive_rna_constraints(X, tid, confidence=sim, passes=2)\n        predictions.append(refined)\n\n    return predictions\n\n# ============================================================\n# Build submission FAST (concat dfs)\n# ============================================================\ndfs = []\nstart_time = time.time()\n\nfor idx, r in enumerate(test_seqs.itertuples(index=False)):\n    if idx % 10 == 0:\n        print(f\"Processing {idx} | {time.time()-start_time:.1f}s\")\n\n    tid = r.target_id\n    seq = r.sequence\n    L = len(seq)\n\n    preds = predict_rna_structures(tid, seq, n_predictions=5)\n\n    data = {\n        \"ID\": [f\"{tid}_{j}\" for j in range(1, L + 1)],\n        \"resname\": list(seq),\n        \"resid\": np.arange(1, L + 1, dtype=np.int32),\n    }\n    for i in range(5):\n        data[f\"x_{i+1}\"] = preds[i][:, 0].astype(np.float32)\n        data[f\"y_{i+1}\"] = preds[i][:, 1].astype(np.float32)\n        data[f\"z_{i+1}\"] = preds[i][:, 2].astype(np.float32)\n\n    dfs.append(pd.DataFrame(data))\n\nsub = pd.concat(dfs, ignore_index=True)\n\ncols = [\"ID\", \"resname\", \"resid\"] + [f\"{c}_{i}\" for i in range(1, 6) for c in [\"x\", \"y\", \"z\"]]\ncoord_cols = [c for c in cols if c.startswith((\"x_\", \"y_\", \"z_\"))]\n\nsub[coord_cols] = sub[coord_cols].clip(-999.999, 9999.999)\nsub[cols].to_csv(\"submission.csv\", index=False)\nprint(\"submission.csv saved\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-05T15:23:21.266324Z","iopub.execute_input":"2026-02-05T15:23:21.266618Z","iopub.status.idle":"2026-02-05T15:24:23.987308Z","shell.execute_reply.started":"2026-02-05T15:23:21.266569Z","shell.execute_reply":"2026-02-05T15:24:23.986688Z"}},"outputs":[],"execution_count":null}]}