{"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":10855324,"sourceType":"datasetVersion","datasetId":6742586},{"sourceId":11118830,"sourceType":"datasetVersion","datasetId":6933267},{"sourceId":14519720,"sourceType":"datasetVersion","datasetId":9271415},{"sourceId":311741,"sourceType":"modelInstanceVersion","modelInstanceId":264400,"modelId":285488}],"dockerImageVersionId":31260,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nfrom scipy.spatial.transform import Rotation as R\nimport random\nfrom Bio import pairwise2\nfrom Bio.Seq import Seq\nfrom scipy.spatial import distance_matrix\nimport warnings\nimport os\nfrom glob import glob\n\nwarnings.filterwarnings('ignore')\n\n# === 1. DATA LOADING ===\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# Load validation if available\ntry:\n    validation_seqs = pd.read_csv(DATA_PATH + 'validation_sequences.csv')\n    validation_labels = pd.read_csv(DATA_PATH + 'validation_labels.csv')\n    print(\"✓ Validation data found and combined with train data.\")\n    combined_seqs = pd.concat([train_seqs, validation_seqs], ignore_index=True)\n    combined_labels = pd.concat([train_labels, validation_labels], ignore_index=True)\nexcept FileNotFoundError:\n    print(\"✗ Validation data not found, using only train data.\")\n    combined_seqs = train_seqs\n    combined_labels = train_labels\n\ndef process_labels(labels_df):\n    coords_dict = {}\n    for id_prefix, group in labels_df.groupby(lambda x: labels_df['ID'][x].rsplit('_', 1)[0]):\n        coords = group.sort_values('resid')[['x_1', 'y_1', 'z_1']].values\n        coords_dict[id_prefix] = coords\n    return coords_dict\n\ncombined_coords_dict = process_labels(combined_labels)\nprint(f\"✓ Loaded {len(combined_coords_dict)} template structures\")\n\n# === 2. MSA DEPTH UTILITY ===\n\ndef get_msa_depth(target_id):\n    msa_path = f'{DATA_PATH}MSA/{target_id}.MSA.fasta'\n    try:\n        with open(msa_path) as f:\n            depth = sum(1 for line in f if line.startswith('>'))\n        return depth\n    except:\n        return 0\n\nprint(\"Computing MSA depths...\")\nmsa_depths = {}\nfor target_id in combined_seqs['target_id'].unique():\n    msa_depths[target_id] = get_msa_depth(target_id)\nprint(f\"✓ Computed MSA depths for {len(msa_depths)} targets\")\n\n# === 3. IMPROVED TEMPLATE SEARCH ===\n\ndef find_similar_sequences(query_seq, train_seqs_df, train_coords_dict, top_n=10):\n    similar_seqs = []\n    query_seq_obj = Seq(query_seq)\n    query_len = len(query_seq)\n\n    for _, row in train_seqs_df.iterrows():\n        target_id, train_seq = row['target_id'], row['sequence']\n        if target_id not in train_coords_dict:\n            continue\n\n        len_diff = abs(len(train_seq) - query_len) / max(len(train_seq), query_len)\n        if len_diff > 0.4:\n            continue\n\n        alignments = pairwise2.align.globalms(\n            query_seq_obj, train_seq, \n            2, -1, -3, -0.1,\n            one_alignment_only=True\n        )\n\n        if alignments:\n            raw_score = alignments[0].score\n            max_possible = 2 * min(query_len, len(train_seq))\n            similarity = raw_score / max_possible\n            length_bonus = 1.0 - len_diff * 0.5\n            msa_bonus = min(1.0 + msa_depths.get(target_id, 0) / 100, 1.3)\n            final_score = similarity * length_bonus * msa_bonus\n\n            similar_seqs.append({\n                'target_id': target_id,\n                'sequence': train_seq,\n                'similarity': similarity,\n                'final_score': final_score,\n                'coords': train_coords_dict[target_id],\n                'msa_depth': msa_depths.get(target_id, 0)\n            })\n\n    similar_seqs.sort(key=lambda x: x['final_score'], reverse=True)\n    return similar_seqs[:top_n]\n\n# === 4. IMPROVED TEMPLATE ADAPTATION ===\n\ndef generate_helix_structure(sequence, seed=None):\n    if seed is not None:\n        np.random.seed(seed)\n    n = len(sequence)\n    coords = np.zeros((n, 3))\n    rise_per_residue = 2.8\n    radius = 10.0\n    residues_per_turn = 11\n    for i in range(n):\n        angle = 2 * np.pi * i / residues_per_turn\n        coords[i] = [radius * np.cos(angle), radius * np.sin(angle), i * rise_per_residue]\n    coords += np.random.normal(0, 0.5, coords.shape)\n    return coords\n\ndef generate_random_structure(sequence, seed=None):\n    if seed is not None:\n        np.random.seed(seed)\n    n = len(sequence)\n    coords = np.zeros((n, 3))\n    for i in range(1, n):\n        theta = np.random.uniform(0, 2*np.pi)\n        phi = np.random.uniform(0, np.pi)\n        r = np.random.uniform(5.5, 6.5)\n        dx = r * np.sin(phi) * np.cos(theta)\n        dy = r * np.sin(phi) * np.sin(theta)\n        dz = r * np.cos(phi)\n        coords[i] = coords[i-1] + [dx, dy, dz]\n    return coords\n\ndef fill_gaps_helix_aware(coords, sequence):\n    C1_DISTANCE = 5.9\n    filled = coords.copy()\n    n = len(filled)\n\n    for i in range(n):\n        if np.isnan(filled[i, 0]):\n            prev_valid = next((j for j in range(i-1, -1, -1) if not np.isnan(filled[j, 0])), -1)\n            next_valid = next((j for j in range(i+1, n) if not np.isnan(filled[j, 0])), -1)\n\n            if prev_valid >= 0 and next_valid >= 0:\n                gap_size = next_valid - prev_valid\n                t = (i - prev_valid) / gap_size\n                base_pos = (1-t) * filled[prev_valid] + t * filled[next_valid]\n                angle = t * np.pi * 2 * (gap_size / 10)\n                radius = 2.0\n                helix_offset = np.array([radius * np.cos(angle), radius * np.sin(angle), 0])\n                filled[i] = base_pos + helix_offset * 0.3\n            elif prev_valid >= 0:\n                if prev_valid > 0:\n                    direction = filled[prev_valid] - filled[prev_valid-1]\n                    direction = direction / (np.linalg.norm(direction) + 1e-10)\n                else:\n                    direction = np.array([1.0, 0.0, 0.0])\n                filled[i] = filled[prev_valid] + direction * C1_DISTANCE\n            elif next_valid >= 0:\n                if next_valid < n - 1:\n                    direction = filled[next_valid] - filled[next_valid+1]\n                    direction = direction / (np.linalg.norm(direction) + 1e-10)\n                else:\n                    direction = np.array([-1.0, 0.0, 0.0])\n                filled[i] = filled[next_valid] + direction * C1_DISTANCE\n            else:\n                filled[i] = np.array([i * C1_DISTANCE * 0.8, 0, 0])\n\n    return np.nan_to_num(filled)\n\ndef adapt_template_to_query(query_seq, template_seq, template_coords):\n    alignments = pairwise2.align.globalms(\n        Seq(query_seq), Seq(template_seq), \n        2, -1, -3, -0.1, \n        one_alignment_only=True\n    )\n\n    if not alignments:\n        return generate_helix_structure(query_seq)\n\n    a_q, a_t = alignments[0].seqA, alignments[0].seqB\n    new_coords = np.full((len(query_seq), 3), np.nan)\n\n    q_idx, t_idx = 0, 0\n    for char_q, char_t in zip(a_q, a_t):\n        if char_q != '-' and char_t != '-':\n            if t_idx < len(template_coords):\n                new_coords[q_idx] = template_coords[t_idx]\n            q_idx += 1\n            t_idx += 1\n        elif char_q != '-':\n            q_idx += 1\n        elif char_t != '-':\n            t_idx += 1\n\n    new_coords = fill_gaps_helix_aware(new_coords, query_seq)\n    return new_coords\n\n# === 5. RNA CONSTRAINTS ===\n\ndef apply_rna_constraints(coordinates, sequence, strength=0.3):\n    refined = coordinates.copy()\n    n = len(sequence)\n    MIN_DIST, MAX_DIST, TARGET_DIST = 5.5, 6.5, 5.9\n\n    for _ in range(3):\n        for i in range(n - 1):\n            dist = np.linalg.norm(refined[i+1] - refined[i])\n            if dist < MIN_DIST or dist > MAX_DIST:\n                direction = (refined[i+1] - refined[i]) / (dist + 1e-10)\n                adjustment = (TARGET_DIST - dist) * strength\n                refined[i] -= direction * adjustment * 0.5\n                refined[i+1] += direction * adjustment * 0.5\n    return refined\n\ndef remove_clashes(coordinates, min_distance=3.5):\n    refined = coordinates.copy()\n    n = len(refined)\n    for _ in range(2):\n        for i in range(n):\n            for j in range(i + 3, n):\n                dist = np.linalg.norm(refined[i] - refined[j])\n                if dist < min_distance:\n                    direction = (refined[j] - refined[i]) / (dist + 1e-10)\n                    push = (min_distance - dist) * 0.5\n                    refined[i] -= direction * push * 0.5\n                    refined[j] += direction * push * 0.5\n    return refined\n\n# === 6. DIVERSE PREDICTION ENSEMBLE ===\n\ndef create_diverse_predictions(sequence, similar_templates, n_predictions=5):\n    predictions = []\n\n    if not similar_templates:\n        for i in range(n_predictions):\n            if i < 2:\n                predictions.append(generate_helix_structure(sequence, seed=i))\n            else:\n                predictions.append(generate_random_structure(sequence, seed=i))\n        return predictions\n\n    # PRED 1: Best template (clean)\n    best = similar_templates[0]\n    pred1 = adapt_template_to_query(sequence, best['sequence'], best['coords'])\n    pred1 = apply_rna_constraints(pred1, sequence, strength=0.3)\n    pred1 = remove_clashes(pred1)\n    predictions.append(pred1)\n\n    # PRED 2: Second best template\n    if len(similar_templates) >= 2:\n        second = similar_templates[1]\n        pred2 = adapt_template_to_query(sequence, second['sequence'], second['coords'])\n        pred2 = apply_rna_constraints(pred2, sequence, strength=0.3)\n        pred2 = remove_clashes(pred2)\n        predictions.append(pred2)\n    else:\n        pred2 = pred1 + np.random.normal(0, 1.0, pred1.shape)\n        pred2 = apply_rna_constraints(pred2, sequence, strength=0.4)\n        predictions.append(pred2)\n\n    # PRED 3: Weighted average of top 3 templates\n    if len(similar_templates) >= 3:\n        weights = np.array([t['final_score'] for t in similar_templates[:3]])\n        weights = weights / weights.sum()\n        templates = [adapt_template_to_query(sequence, t['sequence'], t['coords']) for t in similar_templates[:3]]\n        pred3 = sum(w * t for w, t in zip(weights, templates))\n        pred3 = apply_rna_constraints(pred3, sequence, strength=0.3)\n        pred3 = remove_clashes(pred3)\n        predictions.append(pred3)\n    else:\n        templates = [adapt_template_to_query(sequence, t['sequence'], t['coords']) for t in similar_templates]\n        pred3 = np.mean(templates, axis=0)\n        pred3 = apply_rna_constraints(pred3, sequence, strength=0.3)\n        predictions.append(pred3)\n\n    # PRED 4: Best template + rotation perturbation\n    pred4 = pred1.copy()\n    center = pred4.mean(axis=0)\n    pred4_centered = pred4 - center\n    rotation = R.from_euler('xyz', np.random.uniform(-10, 10, 3), degrees=True)\n    pred4 = rotation.apply(pred4_centered) + center\n    pred4 += np.random.normal(0, 0.3, pred4.shape)\n    pred4 = apply_rna_constraints(pred4, sequence, strength=0.4)\n    predictions.append(pred4)\n\n    # PRED 5: Template + helix hybrid\n    helix = generate_helix_structure(sequence, seed=42)\n    best_similarity = similar_templates[0]['similarity']\n    template_weight = min(best_similarity * 1.2, 0.9)\n    pred5 = template_weight * pred1 + (1 - template_weight) * helix\n    pred5 = apply_rna_constraints(pred5, sequence, strength=0.5)\n    pred5 = remove_clashes(pred5)\n    predictions.append(pred5)\n\n    return predictions[:n_predictions]\n\n# === 7. MAIN PREDICTION FUNCTION ===\n\ndef predict_rna_structures(sequence, target_id, train_seqs_df, train_coords_dict, n_predictions=5):\n    similar_templates = find_similar_sequences(sequence, train_seqs_df, train_coords_dict, top_n=10)\n    predictions = create_diverse_predictions(sequence, similar_templates, n_predictions)\n\n    validated_predictions = []\n    for pred in predictions:\n        if pred.shape[0] != len(sequence):\n            pred = generate_helix_structure(sequence)\n        validated_predictions.append(pred)\n\n    return validated_predictions\n\n# === 8. MAIN LOOP ===\n\nprint(\"\\n\" + \"=\"*50)\nprint(\"Starting predictions...\")\nprint(\"=\"*50)\n\nall_predictions = []\nstart_time = pd.Timestamp.now()\n\nfor idx, row in test_seqs.iterrows():\n    target_id = row['target_id']\n    sequence = row['sequence']\n\n    if idx % 5 == 0:\n        elapsed = (pd.Timestamp.now() - start_time).total_seconds()\n        print(f\"Processing {idx+1}/{len(test_seqs)} | Elapsed: {elapsed:.1f}s | Target: {target_id[:20]}...\")\n\n    preds = predict_rna_structures(sequence, target_id, combined_seqs, combined_coords_dict, n_predictions=5)\n\n    for j in range(len(sequence)):\n        res = {'ID': f\"{target_id}_{j+1}\", 'resname': sequence[j], 'resid': j + 1}\n        for i in range(5):\n            coords = np.clip(preds[i][j], -999.999, 9999.999)\n            res[f'x_{i+1}'] = coords[0]\n            res[f'y_{i+1}'] = coords[1]\n            res[f'z_{i+1}'] = coords[2]\n        all_predictions.append(res)\n\n# === 9. SAVE SUBMISSION ===\n\nsubmission_df = pd.DataFrame(all_predictions)\ncols = ['ID', 'resname', 'resid'] + [f'{c}_{i}' for i in range(1, 6) for c in ['x', 'y', 'z']]\nsubmission_df[cols].to_csv('submission.csv', index=False)\n\ntotal_time = (pd.Timestamp.now() - start_time).total_seconds()\nprint(\"\\n\" + \"=\"*50)\nprint(f\"✓ submission.csv generated!\")\nprint(f\"✓ Total predictions: {len(test_seqs)} sequences\")\nprint(f\"✓ Total rows: {len(submission_df)}\")\nprint(f\"✓ Time elapsed: {total_time:.1f}s\")\nprint(\"=\"*50)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}