{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":118765,"databundleVersionId":15231210}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat":4,"nbformat_minor":4,"cells":[{"cell_type":"markdown","source":"# RNA 3D Folding v2 — Improved TBM\n\nImprovements over v1:\n- Needleman-Wunsch sequence alignment (handles indels properly)\n- Alignment-guided coordinate transfer (no blind interpolation)\n- Top-K template blending (weighted average of best templates)\n- Rotational perturbations for sample diversity","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nfrom scipy.spatial.transform import Rotation\n\nBASE = '/kaggle/input/competitions/stanford-rna-3d-folding-2'\nif not os.path.exists(BASE):\n    BASE = '/kaggle/input/stanford-rna-3d-folding-2'\nprint(f'Data path: {BASE}')\n\nN_SAMPLE = 5\nTOP_K = 3  # number of templates to blend","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load data\ntest_df = pd.read_csv(f'{BASE}/test_sequences.csv')\ntrain_seqs = pd.read_csv(f'{BASE}/train_sequences.csv')\ntrain_labels = pd.read_csv(f'{BASE}/train_labels.csv')\nval_seqs = pd.read_csv(f'{BASE}/validation_sequences.csv')\nval_labels = pd.read_csv(f'{BASE}/validation_labels.csv')\n\nall_seqs = pd.concat([train_seqs, val_seqs], ignore_index=True)\nall_labels = pd.concat([train_labels, val_labels], ignore_index=True)\n\nprint(f'Test targets: {len(test_df)}')\nprint(f'Template pool: {len(all_seqs)} sequences')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Build coord lookup\nall_labels['prefix'] = all_labels['ID'].str.rsplit('_', n=1).str[0]\ncoord_dict = {}\nfor prefix, grp in all_labels.groupby('prefix'):\n    coord_dict[prefix] = grp.sort_values('resid')[['x_1','y_1','z_1']].values\n\nseq_dict = dict(zip(all_seqs['target_id'], all_seqs['sequence']))\nprint(f'Structures loaded: {len(coord_dict)}')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Needleman-Wunsch alignment\ndef nw_align(s1, s2, match=2, mismatch=-1, gap=-2):\n    \"\"\"Needleman-Wunsch global alignment. Returns aligned strings and score.\"\"\"\n    n, m = len(s1), len(s2)\n    # Score matrix\n    dp = np.zeros((n+1, m+1), dtype=np.float32)\n    dp[1:, 0] = np.arange(1, n+1) * gap\n    dp[0, 1:] = np.arange(1, m+1) * gap\n    \n    # Traceback matrix: 0=diag, 1=up, 2=left\n    tb = np.zeros((n+1, m+1), dtype=np.int8)\n    tb[1:, 0] = 1\n    tb[0, 1:] = 2\n    \n    for i in range(1, n+1):\n        for j in range(1, m+1):\n            s = match if s1[i-1] == s2[j-1] else mismatch\n            scores = [dp[i-1, j-1] + s, dp[i-1, j] + gap, dp[i, j-1] + gap]\n            best = int(np.argmax(scores))\n            dp[i, j] = scores[best]\n            tb[i, j] = best\n    \n    # Traceback\n    a1, a2 = [], []\n    i, j = n, m\n    while i > 0 or j > 0:\n        if i > 0 and j > 0 and tb[i, j] == 0:\n            a1.append(s1[i-1]); a2.append(s2[j-1])\n            i -= 1; j -= 1\n        elif i > 0 and tb[i, j] == 1:\n            a1.append(s1[i-1]); a2.append('-')\n            i -= 1\n        else:\n            a1.append('-'); a2.append(s2[j-1])\n            j -= 1\n    \n    return ''.join(reversed(a1)), ''.join(reversed(a2)), float(dp[n, m])\n\n\ndef alignment_score(s1, s2):\n    \"\"\"Fast NW score normalized by max length.\"\"\"\n    _, _, score = nw_align(s1, s2)\n    maxlen = max(len(s1), len(s2))\n    return score / (maxlen * 2) if maxlen > 0 else 0.0  # normalize by max possible\n\nprint('Alignment functions ready')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def kmer_jaccard(s1, s2, k=3):\n    \"\"\"Fast k-mer Jaccard similarity for pre-filtering.\"\"\"\n    if len(s1) < k or len(s2) < k:\n        return 0.0\n    kmers1 = set(s1[i:i+k] for i in range(len(s1)-k+1))\n    kmers2 = set(s2[i:i+k] for i in range(len(s2)-k+1))\n    union = len(kmers1 | kmers2)\n    if union == 0:\n        return 0.0\n    return len(kmers1 & kmers2) / union\n\n\ndef find_top_k_templates(query_seq, seq_dict, coord_dict, k=3, prefilter_n=50):\n    \"\"\"Two-stage template search: fast k-mer filter -> NW alignment on finalists.\"\"\"\n    qlen = len(query_seq)\n    \n    # Stage 1: fast k-mer Jaccard on all candidates\n    kmer_scores = []\n    for tid, tseq in seq_dict.items():\n        if tid not in coord_dict:\n            continue\n        tlen = len(tseq)\n        if abs(tlen - qlen) / max(tlen, qlen) > 0.8:\n            continue\n        jsim = kmer_jaccard(query_seq, tseq)\n        if jsim > 0.01:\n            kmer_scores.append((tid, jsim))\n    \n    kmer_scores.sort(key=lambda x: -x[1])\n    finalists = kmer_scores[:prefilter_n]\n    \n    if not finalists:\n        return []\n    \n    # Stage 2: NW alignment on top candidates only\n    # Cap sequence lengths for NW to avoid timeouts on very long seqs\n    MAX_NW_LEN = 500\n    q_capped = query_seq[:MAX_NW_LEN]\n    \n    candidates = []\n    for tid, _ in finalists:\n        tseq = seq_dict[tid]\n        t_capped = tseq[:MAX_NW_LEN]\n        score = alignment_score(q_capped, t_capped)\n        candidates.append((tid, score))\n    \n    candidates.sort(key=lambda x: -x[1])\n    return candidates[:k]\n\nprint('Template search ready')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def alignment_guided_coords(query_seq, template_seq, template_coords):\n    \"\"\"Transfer coordinates using NW alignment.\n    Aligned positions get template coords directly.\n    Gaps are filled by interpolation between nearest assigned positions.\n    \"\"\"\n    a1, a2, _ = nw_align(query_seq, template_seq)\n    \n    qlen = len(query_seq)\n    coords = np.full((qlen, 3), np.nan)\n    \n    qi, ti = 0, 0\n    for c1, c2 in zip(a1, a2):\n        if c1 != '-' and c2 != '-':\n            # Aligned position: copy template coords\n            if ti < len(template_coords):\n                coords[qi] = template_coords[ti]\n            qi += 1; ti += 1\n        elif c1 != '-' and c2 == '-':\n            # Query has residue, template has gap — will interpolate\n            qi += 1\n        else:\n            # Template has residue, query has gap — skip template residue\n            ti += 1\n    \n    # Interpolate NaN positions\n    for d in range(3):\n        col = coords[:, d]\n        valid = ~np.isnan(col)\n        if valid.sum() >= 2:\n            coords[:, d] = np.interp(\n                np.arange(qlen),\n                np.where(valid)[0],\n                col[valid]\n            )\n        elif valid.sum() == 1:\n            coords[:, d] = col[valid][0]\n        else:\n            coords[:, d] = np.linspace(0, qlen * 2.5, qlen)\n    \n    return coords\n\n\ndef blend_templates(query_seq, templates, seq_dict, coord_dict):\n    \"\"\"Weighted average of alignment-guided coords from multiple templates.\"\"\"\n    if not templates:\n        return denovo_coords(len(query_seq))\n    \n    qlen = len(query_seq)\n    weighted_coords = np.zeros((qlen, 3))\n    total_weight = 0\n    \n    for tid, score in templates:\n        w = max(score, 0.01) ** 2  # square weighting to favor better matches\n        tcoords = alignment_guided_coords(query_seq, seq_dict[tid], coord_dict[tid])\n        \n        # Superimpose subsequent templates onto the first using centroid alignment\n        if total_weight > 0:\n            ref = weighted_coords / total_weight\n            # Translate to match centroids\n            tcoords = tcoords - tcoords.mean(axis=0) + ref.mean(axis=0)\n        \n        weighted_coords += w * tcoords\n        total_weight += w\n    \n    return weighted_coords / total_weight\n\n\ndef denovo_coords(seq_len):\n    \"\"\"Fallback: A-form helix.\"\"\"\n    coords = np.zeros((seq_len, 3))\n    for i in range(seq_len):\n        ang = i * 0.6\n        coords[i] = [10*np.cos(ang), 10*np.sin(ang), i*2.5]\n    return coords\n\n\ndef perturb_coords(coords, scale=0.5, seed=None):\n    \"\"\"Apply a small random rotation + jitter for sample diversity.\"\"\"\n    rng = np.random.RandomState(seed)\n    # Small random rotation (up to ~5 degrees)\n    angle = rng.normal(0, scale * 0.05)  # radians\n    axis = rng.randn(3)\n    axis /= np.linalg.norm(axis) + 1e-8\n    rot = Rotation.from_rotvec(angle * axis)\n    \n    centroid = coords.mean(axis=0)\n    rotated = rot.apply(coords - centroid) + centroid\n    \n    # Small per-residue jitter\n    jitter = rng.normal(0, scale * 0.3, coords.shape)\n    return rotated + jitter\n\nprint('Coordinate functions ready')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Generate predictions\nrows = []\nfor idx, row in test_df.iterrows():\n    tid = row['target_id']\n    seq = row['sequence']\n    slen = len(seq)\n    \n    templates = find_top_k_templates(seq, seq_dict, coord_dict, k=TOP_K)\n    \n    if templates and templates[0][1] > 0.05:\n        coords = blend_templates(seq, templates, seq_dict, coord_dict)\n        t_info = ', '.join(f'{t[0]}({t[1]:.3f})' for t in templates)\n        print(f'{tid} ({slen}nt): templates=[{t_info}]')\n    else:\n        coords = denovo_coords(slen)\n        print(f'{tid} ({slen}nt): de-novo fallback')\n    \n    for i in range(slen):\n        r = {'ID': f'{tid}_{i+1}', 'resname': seq[i], 'resid': i+1}\n        for s in range(1, N_SAMPLE+1):\n            if s == 1:\n                c = coords[i]\n            else:\n                pc = perturb_coords(coords, scale=0.3 * s, seed=42 + s * 1000 + i)\n                c = pc[i]\n            r[f'x_{s}'] = float(np.clip(c[0], -999.999, 9999.999))\n            r[f'y_{s}'] = float(np.clip(c[1], -999.999, 9999.999))\n            r[f'z_{s}'] = float(np.clip(c[2], -999.999, 9999.999))\n        rows.append(r)\n\nprint(f'\\nTotal rows: {len(rows)}')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Save submission\nsub = pd.DataFrame(rows)\ncols = ['ID', 'resname', 'resid'] + [f'{c}_{i}' for i in range(1, N_SAMPLE+1) for c in ['x','y','z']]\nsub[cols].to_csv('/kaggle/working/submission.csv', index=False)\nprint('Saved submission.csv')\nprint(sub[cols].head())\nprint(f'Shape: {sub[cols].shape}')","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}