{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":118765,"databundleVersionId":15231210,"sourceType":"competition"},{"sourceId":14604295,"sourceType":"datasetVersion","datasetId":9328538}],"dockerImageVersionId":31259,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ============================================================\n# RNA 3D FOLDING – FAST TEMPLATE BASELINE (MSA + per‑chain)\n# CPU‑only, deterministic, no training.\n# Output: submission.csv with 5 structures per target.\n# ============================================================\n\nimport sys, os, subprocess, time, warnings, hashlib, zlib\nfrom collections import defaultdict\nos.environ[\"PYTHONHASHSEED\"] = \"0\"\nwarnings.filterwarnings(\"ignore\")\nstart = time.time()\n\n# --- Install Biopython wheel (Kaggle offline) ---\nsubprocess.check_call([\n    sys.executable, \"-m\", \"pip\", \"install\", \"--no-index\",\n    \"/kaggle/input/biopython-cp312/biopython-1.86-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl\"\n])\n\nimport numpy as np\nimport pandas as pd\nfrom Bio.Align import PairwiseAligner\nfrom Bio import SeqIO\n\n# ========== 1. LOAD DATA ==========\nDATA_PATH = \"/kaggle/input/stanford-rna-3d-folding-2/\"\nsys.path.append(os.path.join(DATA_PATH, \"extra\"))\ntry:\n    from parse_fasta_py import parse_fasta\nexcept:\n    def parse_fasta(fasta_str):\n        d, h = {}, None\n        for line in fasta_str.splitlines():\n            line = line.strip()\n            if line.startswith('>'):\n                h = line[1:].split()[0]\n                d[h] = ''\n            elif h:\n                d[h] += line.replace(' ', '')\n        return d\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# ---------- Per‑chain coordinates from training labels ----------\ndef process_labels_per_chain(labels_df):\n    target_dict = {}\n    for tid, group in labels_df.groupby(labels_df['ID'].str.rsplit('_', n=1).str[0]):\n        group = group.copy()\n        group['chain'] = group['chain'].astype(str)          # ← FIX: int→str\n        group['copy']  = group['copy'].astype(int)\n        chain_key = group['chain'] + '_' + group['copy'].astype(str)\n        chain_dict = {}\n        for key, sub in group.groupby(chain_key):\n            sub = sub.sort_values('resid')\n            seq = ''.join(sub['resname'].values)\n            coords = sub[['x_1','y_1','z_1']].values.astype(np.float32)\n            chain_dict[key] = (seq, coords)\n        target_dict[tid] = chain_dict\n    return target_dict\n\ntrain_chain_coords = process_labels_per_chain(train_labels)\n\n# ---------- Build template library (ONLY chains with coordinates) ----------\ntemplate_db = []   # list of dicts: id, seq, coords, length, kmer_set\nkmer_len = 3\n\ndef kmer_set(seq):\n    return {seq[i:i+kmer_len] for i in range(len(seq)-kmer_len+1)}\n\nfor tid, chain_dict in train_chain_coords.items():\n    date_row = train_seqs[train_seqs['target_id'] == tid]\n    date = date_row['temporal_cutoff'].values[0] if not date_row.empty else '1900-01-01'\n    for chain_key, (seq, coords) in chain_dict.items():\n        template_db.append({\n            'id': f\"{tid}_{chain_key}\",\n            'seq': seq,\n            'len': len(seq),\n            'coords': coords,\n            'date': date,\n            'kmers': kmer_set(seq)\n        })\n\nprint(f\"Template DB: {len(template_db)} chains with coordinates\")\n\n# ---------- Stoichiometry segment maps for test ----------\ndef parse_stoich(stoich_str):\n    if pd.isna(stoich_str) or not str(stoich_str).strip():\n        return []\n    parts = [p.split(':') for p in str(stoich_str).split(';')]\n    return [(p[0].strip(), int(p[1].strip())) for p in parts if len(p)==2]\n\ndef get_chain_segments(row):\n    seq_all = row['sequence']\n    stoich = parse_stoich(row.get('stoichiometry',''))\n    if not stoich:\n        return [(0, len(seq_all), 'A', 1)]\n    try:\n        fasta_dict = parse_fasta(row.get('all_sequences',''))\n    except:\n        fasta_dict = {}\n    segments, pos = [], 0\n    for chain_id, copies in stoich:\n        chain_seq = fasta_dict.get(chain_id, '')\n        if not chain_seq:\n            total_copies = sum(c for _,c in stoich)\n            guess_len = len(seq_all)//total_copies if total_copies else 1\n            chain_seq = 'N'*guess_len\n        for copy_idx in range(1, copies+1):\n            L = len(chain_seq)\n            segments.append((pos, pos+L, chain_id, copy_idx))\n            pos += L\n    return segments\n\ntest_seg_map = {}\nfor _, row in test_seqs.iterrows():\n    test_seg_map[row['target_id']] = get_chain_segments(row)\n\n# ========== 2. ALIGNMENT SETUP ==========\naligner = PairwiseAligner()\naligner.mode = 'global'\naligner.match_score = 2.0\naligner.mismatch_score = -1.5\naligner.open_gap_score = -8.0\naligner.extend_gap_score = -0.4\nfor term in ['query_left','query_right','target_left','target_right']:\n    setattr(aligner, f'{term}_open_gap_score', -8.0)\n    setattr(aligner, f'{term}_extend_gap_score', -0.4)\n\n# ========== 3. FAST TEMPLATE SEARCH PER CHAIN ==========\ndef find_templates_for_chain(chain_seq, max_candidates=20):\n    \"\"\"Use k‑mer overlap + length filter, then full DP on top candidates.\"\"\"\n    L = len(chain_seq)\n    min_len, max_len = int(L*0.7), int(L*1.3)\n    candidates = []\n    chain_kmers = kmer_set(chain_seq)\n    \n    # 1st pass: cheap k‑mer Jaccard\n    for t in template_db:\n        if not (min_len <= t['len'] <= max_len):\n            continue\n        intersect = len(chain_kmers & t['kmers'])\n        union = len(chain_kmers) + len(t['kmers']) - intersect\n        if union == 0:\n            jaccard = 0\n        else:\n            jaccard = intersect / union\n        if jaccard > 0.15:                     # heuristic threshold\n            candidates.append((t, jaccard))\n    \n    # sort by k‑mer similarity, take top N for expensive DP\n    candidates.sort(key=lambda x: -x[1])\n    top = []\n    for t, _ in candidates[:max_candidates*2]:   # examine more, keep best\n        score = aligner.score(chain_seq, t['seq']) / (2 * min(L, t['len']))\n        top.append((t, score))\n    top.sort(key=lambda x: -x[1])\n    return [(t['seq'], t['coords'], score) for t, score in top[:max_candidates]]\n\n# ========== 4. COORDINATE TRANSFER ==========\ndef adapt_coords(query_seq, template_seq, template_coords):\n    \"\"\"Global alignment → coordinate transfer + interpolation.\"\"\"\n    if len(template_seq) != len(template_coords):\n        raise ValueError(\"Length mismatch\")\n    aln = next(aligner.align(query_seq, template_seq))\n    new_coords = np.full((len(query_seq), 3), np.nan, dtype=np.float32)\n    for (qs,qe), (ts,te) in zip(*aln.aligned):\n        if qe-qs == te-ts:\n            new_coords[qs:qe] = template_coords[ts:te]\n    # fill gaps\n    for i in range(len(new_coords)):\n        if np.isnan(new_coords[i,0]):\n            prev = next((j for j in range(i-1,-1,-1) if not np.isnan(new_coords[j,0])), -1)\n            nxt  = next((j for j in range(i+1,len(new_coords)) if not np.isnan(new_coords[j,0])), -1)\n            if prev>=0 and nxt>=0:\n                w = (i-prev)/(nxt-prev)\n                new_coords[i] = (1-w)*new_coords[prev] + w*new_coords[nxt]\n            elif prev>=0:\n                new_coords[i] = new_coords[prev] + [3.0,0,0]\n            elif nxt>=0:\n                new_coords[i] = new_coords[nxt] + [-3.0,0,0]\n            else:\n                new_coords[i] = [i*3.0,0,0]\n    return np.nan_to_num(new_coords)\n\ndef consensus_coordinates(adapted_list, weights=None):\n    if not adapted_list:\n        return None\n    if weights is None:\n        weights = np.ones(len(adapted_list))\n    weights = np.array(weights)/np.sum(weights)\n    result = np.zeros_like(adapted_list[0])\n    for w, c in zip(weights, adapted_list):\n        result += w * c\n    return result\n\n# ========== 5. REFINEMENT (unchanged, optimised) ==========\ndef refine_structure(coords, segments, confidence=0.8, passes=2):\n    strength = 0.7 * (1.0 - min(confidence, 0.9))\n    strength = max(strength, 0.05)\n    coords = coords.copy()\n    for _ in range(passes):\n        for s,e,_,_ in segments:\n            X = coords[s:e].copy()\n            L = e-s\n            if L<4: continue\n            # bond length\n            d = X[1:]-X[:-1]\n            dist = np.linalg.norm(d, axis=1, keepdims=True)+1e-6\n            scale = (5.95 - dist)/dist\n            adj = d * scale * 0.25 * strength\n            X[:-1] -= adj\n            X[1:]  += adj\n            # i,i+2\n            if L>2:\n                d2 = X[2:]-X[:-2]\n                dist2 = np.linalg.norm(d2, axis=1, keepdims=True)+1e-6\n                scale2 = (10.2 - dist2)/dist2\n                adj2 = d2 * scale2 * 0.12 * strength\n                X[:-2] -= adj2\n                X[2:]  += adj2\n            # smoothing\n            if L>3:\n                lap = 0.5*(X[:-2]+X[2:]) - X[1:-1]\n                X[1:-1] += lap * 0.08 * strength\n            coords[s:e] = X\n    return coords\n\n# ========== 6. DIVERSITY GENERATORS ==========\ndef rotate_chain(coords, segment, rng, max_angle=15):\n    s,e,_,_ = segment\n    if e-s < 20: return coords\n    pivot = s + rng.integers(10, e-s-10)\n    axis = rng.normal(size=3); axis /= (np.linalg.norm(axis)+1e-8)\n    ang = np.deg2rad(rng.uniform(-max_angle,max_angle))\n    c,s_ = np.cos(ang), np.sin(ang)\n    K = np.array([[0,-axis[2],axis[1]],[axis[2],0,-axis[0]],[-axis[1],axis[0],0]])\n    R = np.eye(3) + s_*K + (1-c)*(K@K)\n    X = coords.copy()\n    center = X[pivot].copy()\n    X[pivot+1:e] = (X[pivot+1:e] - center) @ R.T + center\n    return X\n\ndef jitter_chains(coords, segments, rng, trans_max=1.2, angle_max=8):\n    X = coords.copy()\n    for s,e,_,_ in segments:\n        if e-s < 10: continue\n        axis = rng.normal(size=3); axis /= (np.linalg.norm(axis)+1e-8)\n        ang = np.deg2rad(rng.uniform(-angle_max,angle_max))\n        c,s_ = np.cos(ang), np.sin(ang)\n        K = np.array([[0,-axis[2],axis[1]],[axis[2],0,-axis[0]],[-axis[1],axis[0],0]])\n        R = np.eye(3) + s_*K + (1-c)*(K@K)\n        shift = rng.uniform(-trans_max, trans_max, 3)\n        center = X[s:e].mean(axis=0)\n        X[s:e] = (X[s:e]-center) @ R.T + center + shift\n    return X\n\ndef gaussian_noise(coords, rng, sigma=0.2):\n    return coords + rng.normal(0, sigma, coords.shape)\n\n# ========== 7. DETERMINISTIC SEED ==========\ndef stable_seed(s):\n    return zlib.adler32(s.encode()) & 0xffffffff\n\n# ========== 8. MAIN PREDICTION LOOP ==========\nprint(\"Starting per‑chain prediction...\")\nall_rows = []\n\nfor idx, row in test_seqs.iterrows():\n    if idx % 5 == 0:\n        print(f\"Processing {idx}/{len(test_seqs)} | {time.time()-start:.1f}s\")\n    tid = row['target_id']\n    full_seq = row['sequence']\n    segments = test_seg_map.get(tid, [(0,len(full_seq),'A',1)])\n    \n    # --- Build each chain independently ---\n    chain_models = {}\n    for s,e,chain_id,copy in segments:\n        chain_seq = full_seq[s:e]\n        templates = find_templates_for_chain(chain_seq, max_candidates=10)\n        if not templates:\n            # fallback: straight line\n            fallback = np.zeros((e-s,3))\n            for i in range(1,e-s):\n                fallback[i] = fallback[i-1] + [5.95,0,0]\n            chain_models[(s,e)] = fallback\n        else:\n            adapted_list = []\n            scores = []\n            for templ_seq, templ_coords, score in templates[:3]:  # top 3\n                adapted = adapt_coords(chain_seq, templ_seq, templ_coords)\n                adapted_list.append(adapted)\n                scores.append(score)\n            consensus = consensus_coordinates(adapted_list, scores)\n            chain_models[(s,e)] = consensus\n    \n    # Stitch chains\n    base_coords = np.zeros((len(full_seq),3))\n    for (s,e), coords in chain_models.items():\n        base_coords[s:e] = coords\n    \n    # --- 5 predictions with diversity ---\n    preds = []\n    for i in range(5):\n        seed = stable_seed(tid + f\"_{i}\")\n        rng = np.random.default_rng(seed)\n        if i == 0:\n            X = base_coords.copy()\n        elif i == 1:\n            X = jitter_chains(base_coords, segments, rng, trans_max=0.8, angle_max=6)\n        elif i == 2:\n            longest = max(segments, key=lambda x: x[1]-x[0])\n            X = rotate_chain(base_coords, longest, rng, max_angle=18)\n        elif i == 3:\n            X = gaussian_noise(base_coords, rng, sigma=0.3)\n        else:\n            X = jitter_chains(base_coords, segments, rng, trans_max=1.2, angle_max=10)\n        refined = refine_structure(X, segments, confidence=0.7, passes=2)\n        preds.append(refined)\n    \n    # --- Write submission rows ---\n    L = len(full_seq)\n    for res_i in range(L):\n        row_dict = {\n            'ID': f\"{tid}_{res_i+1}\",\n            'resname': full_seq[res_i],\n            'resid': res_i+1\n        }\n        for p in range(5):\n            c = preds[p][res_i]\n            row_dict[f'x_{p+1}'] = c[0]\n            row_dict[f'y_{p+1}'] = c[1]\n            row_dict[f'z_{p+1}'] = c[2]\n        all_rows.append(row_dict)\n\n# ========== 9. SAVE SUBMISSION ==========\nsub_df = pd.DataFrame(all_rows)\ncoord_cols = [f'{c}_{i}' for i in range(1,6) for c in ['x','y','z']]\nsub_df[coord_cols] = sub_df[coord_cols].clip(-999.999, 9999.999)\nsub_df = sub_df[['ID','resname','resid'] + coord_cols]\nsub_df.to_csv('submission.csv', index=False)\nprint(f\"Submission saved! Total time: {time.time()-start:.1f}s\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T03:39:55.065396Z","iopub.execute_input":"2026-02-08T03:39:55.065732Z","iopub.status.idle":"2026-02-08T03:43:25.947347Z","shell.execute_reply.started":"2026-02-08T03:39:55.065703Z","shell.execute_reply":"2026-02-08T03:43:25.946312Z"}},"outputs":[],"execution_count":null}]}