{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":118765,"databundleVersionId":15231210,"sourceType":"competition"}],"dockerImageVersionId":31259,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# RNA 3D Folding - Comprehensive Data Cleaning\n#\nThis script covers ALL possible data cleaning methods for the RNA structure prediction task.\n#\n## Table of Contents\n1. **Data Loading & Initial Exploration**\n2. **Missing Value Analysis**\n3. **Coordinate Quality Assessment**\n4. **Sequence Quality Assessment**\n5. **Outlier Detection & Removal**\n6. **Bond Length Validation & Correction**\n7. **Duplicate/Redundancy Removal**\n8. **Structure Quality Metrics**\n9. **Secondary Structure Validation (RNAfold)**\n10. **Resolution-based Filtering**\n11. **Temporal Validation**\n12. **Cross-validation with PDB**\n13. **Ensemble Averaging**\n14. **Chain Break Detection**\n15. **Energy Minimization**\n16. **Final Cleaned Dataset Export**","metadata":{"_uuid":"1a191a6a-366f-4331-baa7-0129697c15d4","_cell_guid":"031a449d-cf84-4b90-8735-2358424b658f","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom collections import Counter, defaultdict\nimport warnings\nimport os\nimport subprocess\nimport re\nfrom datetime import datetime\nfrom typing import Dict, List, Tuple, Optional, Set\nimport json\n\nwarnings.filterwarnings('ignore')\n\n# Set display options\npd.set_option('display.max_columns', None)\npd.set_option('display.max_rows', 100)\npd.set_option('display.width', None)\n\n# Update this path to match your data location\nDATA_PATH = '/kaggle/input/stanford-rna-3d-folding-2/'\nOUTPUT_PATH = './'  # Output directory for cleaned files\n\n# Create output directory if it doesn't exist\nos.makedirs(OUTPUT_PATH, exist_ok=True)","metadata":{"_uuid":"d33be624-3723-4902-9368-5b18f550d8dd","_cell_guid":"db7f3a9e-fa36-4c94-a58d-c71c1a2c51a9","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 1. Data Loading & Initial Exploration","metadata":{"_uuid":"7487b34b-b51b-416d-b408-0d31b3ea6e7a","_cell_guid":"1d09f74b-ce8e-4fdf-8cfb-7d554d8a60b3","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"# Load the datasets\nprint(\"Loading datasets...\")\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\nprint(f\"\\n{'='*60}\")\nprint(\"DATASET SHAPES\")\nprint(f\"{'='*60}\")\nprint(f\"train_sequences: {train_seqs.shape}\")\nprint(f\"test_sequences:  {test_seqs.shape}\")\nprint(f\"train_labels:    {train_labels.shape}\")","metadata":{"_uuid":"590cdb13-cce0-4383-a5ea-b268d575c9fe","_cell_guid":"48e8ff1b-aba9-4211-aadc-a62b7eb56e18","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Process labels to get per-structure coordinates\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):\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\nprint(f\"\\n{'='*60}\")\nprint(\"INITIAL STATISTICS\")\nprint(f\"{'='*60}\")\nprint(f\"Number of unique structures: {len(train_coords_dict)}\")\nprint(f\"Number of sequences: {len(train_seqs)}\")","metadata":{"_uuid":"32b727a3-3076-4d02-b2eb-82a08c9e4127","_cell_guid":"c0b8a4f3-e45e-4121-8727-b8f9dc641998","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 2-8. Core Cleaning Functions (Previous Implementation)","metadata":{"_uuid":"00e5d424-0970-4c95-a712-c38e85667c69","_cell_guid":"5907b7d5-4ae9-4487-9e89-caef3cc44802","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"def fix_bond_lengths(coords, target_bond=5.95, max_iterations=20, tolerance=0.1):\n    \"\"\"\n    Fix bond lengths using iterative relaxation.\n    \"\"\"\n    coords = coords.copy().astype(float)\n    n = len(coords)\n\n    if n < 2:\n        return coords\n\n    for iteration in range(max_iterations):\n        diffs = coords[1:] - coords[:-1]\n        distances = np.linalg.norm(diffs, axis=1) + 1e-8\n        errors = target_bond - distances\n        max_correction = np.max(np.abs(errors))\n\n        if max_correction < tolerance:\n            break\n\n        scale = errors / distances\n        adjustments = diffs * scale[:, np.newaxis] * 0.25\n        coords[:-1] -= adjustments\n        coords[1:] += adjustments\n\n    return coords","metadata":{"_uuid":"db02d050-eeb6-422d-8276-6b1e69265a4b","_cell_guid":"06a7fc02-ab98-4de4-a0bd-bb6237c1add1","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 9. Secondary Structure Validation (RNAfold)","metadata":{"_uuid":"d049e36c-e6f0-4e3e-9ec9-9d23dfad2da0","_cell_guid":"1258de50-1314-4c21-a374-b7c1c828f629","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"class SecondaryStructureValidator:\n    \"\"\"\n    Validate RNA structures using secondary structure prediction.\n    Uses RNAfold from ViennaRNA package if available, otherwise uses simple heuristics.\n    \"\"\"\n\n    def __init__(self):\n        self.rnafold_available = self._check_rnafold()\n        if not self.rnafold_available:\n            print(\"Warning: RNAfold not found. Using heuristic secondary structure validation.\")\n\n    def _check_rnafold(self) -> bool:\n        \"\"\"Check if RNAfold is available.\"\"\"\n        try:\n            result = subprocess.run(['RNAfold', '--version'],\n                                    capture_output=True, text=True, timeout=5)\n            return result.returncode == 0\n        except (FileNotFoundError, subprocess.TimeoutExpired):\n            return False\n\n    def predict_secondary_structure(self, sequence: str) -> Tuple[str, float]:\n        \"\"\"\n        Predict secondary structure using RNAfold or heuristics.\n        Returns (dot-bracket notation, MFE)\n        \"\"\"\n        if self.rnafold_available:\n            return self._rnafold_predict(sequence)\n        else:\n            return self._heuristic_predict(sequence)\n\n    def _rnafold_predict(self, sequence: str) -> Tuple[str, float]:\n        \"\"\"Use RNAfold for prediction.\"\"\"\n        try:\n            result = subprocess.run(\n                ['RNAfold', '--noPS'],\n                input=sequence,\n                capture_output=True,\n                text=True,\n                timeout=30\n            )\n            lines = result.stdout.strip().split('\\n')\n            if len(lines) >= 2:\n                # Parse dot-bracket and MFE\n                structure_line = lines[1]\n                match = re.search(r'([.()]+)\\s*\\(\\s*(-?\\d+\\.?\\d*)\\s*\\)', structure_line)\n                if match:\n                    return match.group(1), float(match.group(2))\n            return '.' * len(sequence), 0.0\n        except Exception:\n            return '.' * len(sequence), 0.0\n\n    def _heuristic_predict(self, sequence: str) -> Tuple[str, float]:\n        \"\"\"\n        Simple heuristic secondary structure prediction.\n        Identifies potential base pairs based on complementarity.\n        \"\"\"\n        n = len(sequence)\n        structure = ['.'] * n\n        pairs = []\n\n        # Watson-Crick and wobble pairs\n        valid_pairs = {('A', 'U'), ('U', 'A'), ('G', 'C'), ('C', 'G'), ('G', 'U'), ('U', 'G')}\n\n        # Find potential hairpin loops (minimum 4 nt loop)\n        min_loop = 4\n        for i in range(n):\n            for j in range(i + min_loop + 2, n):\n                if (sequence[i], sequence[j]) in valid_pairs:\n                    # Check if this pair doesn't conflict with existing pairs\n                    conflict = False\n                    for pi, pj in pairs:\n                        if (i < pi < j < pj) or (pi < i < pj < j):\n                            conflict = True\n                            break\n                    if not conflict:\n                        pairs.append((i, j))\n                        structure[i] = '('\n                        structure[j] = ')'\n\n        # Estimate MFE (rough approximation)\n        mfe = -len(pairs) * 2.0  # ~2 kcal/mol per base pair\n\n        return ''.join(structure), mfe\n\n    def validate_3d_vs_2d(self, coords: np.ndarray, sequence: str,\n                          distance_threshold: float = 15.0) -> Dict:\n        \"\"\"\n        Validate 3D structure against predicted secondary structure.\n\n        Base pairs in secondary structure should be close in 3D space.\n        \"\"\"\n        structure, mfe = self.predict_secondary_structure(sequence)\n\n        # Find base pairs from structure\n        stack = []\n        pairs = []\n        for i, char in enumerate(structure):\n            if char == '(':\n                stack.append(i)\n            elif char == ')' and stack:\n                j = stack.pop()\n                pairs.append((j, i))\n\n        if not pairs or len(coords) < 2:\n            return {\n                'structure': structure,\n                'mfe': mfe,\n                'n_pairs': 0,\n                'consistent_pairs': 0,\n                'consistency_score': 1.0,\n                'violations': []\n            }\n\n        # Check distances for base pairs\n        consistent = 0\n        violations = []\n\n        for i, j in pairs:\n            if i < len(coords) and j < len(coords):\n                dist = np.linalg.norm(coords[i] - coords[j])\n                if dist < distance_threshold:\n                    consistent += 1\n                else:\n                    violations.append((i, j, dist))\n\n        consistency_score = consistent / len(pairs) if pairs else 1.0\n\n        return {\n            'structure': structure,\n            'mfe': mfe,\n            'n_pairs': len(pairs),\n            'consistent_pairs': consistent,\n            'consistency_score': consistency_score,\n            'violations': violations\n        }","metadata":{"_uuid":"c8dff310-05c9-4d93-8504-17eb57f9f1af","_cell_guid":"3c2a1b12-64b6-47be-9e76-17850992cc47","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Initialize secondary structure validator\nss_validator = SecondaryStructureValidator()\n\ndef validate_secondary_structures(seqs_df: pd.DataFrame, coords_dict: Dict,\n                                  sample_size: int = 100) -> pd.DataFrame:\n    \"\"\"\n    Validate secondary structures for a sample of structures.\n    \"\"\"\n    print(\"\\n=== Secondary Structure Validation ===\")\n\n    results = []\n    sample_ids = list(coords_dict.keys())[:sample_size]\n\n    for tid in sample_ids:\n        seq_row = seqs_df[seqs_df['target_id'] == tid]\n        if len(seq_row) == 0:\n            continue\n\n        sequence = seq_row.iloc[0]['sequence']\n        coords = coords_dict[tid]\n\n        validation = ss_validator.validate_3d_vs_2d(coords, sequence)\n        validation['target_id'] = tid\n        results.append(validation)\n\n    results_df = pd.DataFrame(results)\n\n    print(f\"  Validated {len(results_df)} structures\")\n    print(f\"  Mean consistency score: {results_df['consistency_score'].mean():.3f}\")\n    print(f\"  Structures with >80% consistency: {(results_df['consistency_score'] > 0.8).sum()}\")\n\n    return results_df","metadata":{"_uuid":"3cd01de4-d986-45bf-925a-a75ed1b2355c","_cell_guid":"8d7dd64a-3c75-404b-bdbd-f3fb15fa74f9","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 10. Resolution-based Filtering","metadata":{"_uuid":"f1aa7839-1068-4ef7-9b77-1fdc962c0026","_cell_guid":"a6e8d2b8-39a0-44c7-a166-edb7280f37fd","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"class ResolutionFilter:\n    \"\"\"\n    Filter structures based on resolution/quality metrics.\n    Since we don't have explicit resolution data, we estimate quality from coordinates.\n    \"\"\"\n\n    def __init__(self):\n        self.quality_metrics = {}\n\n    def estimate_resolution_proxy(self, coords: np.ndarray) -> Dict:\n        \"\"\"\n        Estimate resolution-like quality metrics from coordinates.\n\n        Lower values indicate better quality (like X-ray resolution).\n        \"\"\"\n        if coords is None or len(coords) < 3:\n            return {'resolution_proxy': np.inf, 'quality_class': 'poor'}\n\n        valid_mask = ~(np.isnan(coords).any(axis=1) | np.isinf(coords).any(axis=1))\n        if valid_mask.sum() < 3:\n            return {'resolution_proxy': np.inf, 'quality_class': 'poor'}\n\n        valid_coords = coords[valid_mask]\n\n        # Compute quality metrics\n\n        # 1. Bond length consistency (lower std = better)\n        diffs = valid_coords[1:] - valid_coords[:-1]\n        bond_lengths = np.linalg.norm(diffs, axis=1)\n        bond_std = np.std(bond_lengths)\n\n        # 2. Bond angle consistency\n        if len(valid_coords) >= 3:\n            v1 = valid_coords[:-2] - valid_coords[1:-1]\n            v2 = valid_coords[2:] - valid_coords[1:-1]\n            cos_angles = np.sum(v1 * v2, axis=1) / (\n                np.linalg.norm(v1, axis=1) * np.linalg.norm(v2, axis=1) + 1e-8\n            )\n            angles = np.arccos(np.clip(cos_angles, -1, 1))\n            angle_std = np.std(angles)\n        else:\n            angle_std = np.inf\n\n        # 3. Local density consistency\n        centroid = np.mean(valid_coords, axis=0)\n        distances_from_centroid = np.linalg.norm(valid_coords - centroid, axis=1)\n        density_cv = np.std(distances_from_centroid) / (np.mean(distances_from_centroid) + 1e-8)\n\n        # Combine into resolution proxy (lower is better)\n        resolution_proxy = bond_std + angle_std * 0.5 + density_cv * 0.3\n\n        # Classify quality\n        if resolution_proxy < 0.5:\n            quality_class = 'excellent'\n        elif resolution_proxy < 1.0:\n            quality_class = 'good'\n        elif resolution_proxy < 2.0:\n            quality_class = 'moderate'\n        else:\n            quality_class = 'poor'\n\n        return {\n            'resolution_proxy': resolution_proxy,\n            'bond_std': bond_std,\n            'angle_std': angle_std,\n            'density_cv': density_cv,\n            'quality_class': quality_class\n        }\n\n    def filter_by_quality(self, coords_dict: Dict,\n                          max_resolution_proxy: float = 2.0) -> Tuple[Dict, List]:\n        \"\"\"\n        Filter structures by estimated quality.\n        \"\"\"\n        print(f\"\\n=== Resolution-based Filtering (threshold={max_resolution_proxy}) ===\")\n\n        filtered = {}\n        removed = []\n        quality_data = []\n\n        for tid, coords in coords_dict.items():\n            metrics = self.estimate_resolution_proxy(coords)\n            metrics['target_id'] = tid\n            quality_data.append(metrics)\n\n            if metrics['resolution_proxy'] <= max_resolution_proxy:\n                filtered[tid] = coords\n            else:\n                removed.append(tid)\n\n        self.quality_metrics = pd.DataFrame(quality_data)\n\n        print(f\"  Removed {len(removed)} low-quality structures\")\n        print(f\"  Remaining: {len(filtered)} structures\")\n        print(f\"  Quality distribution:\")\n        print(self.quality_metrics['quality_class'].value_counts())\n\n        return filtered, removed","metadata":{"_uuid":"c90fe859-cf48-4dd9-af67-949b1731d25f","_cell_guid":"375a0c29-f4b9-4021-b85a-311174d1a119","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 11. Temporal Validation","metadata":{"_uuid":"2a977b78-ac6a-47c1-b1c4-9113baf10753","_cell_guid":"82049be8-5104-4351-8043-e358d72facc9","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"class TemporalValidator:\n    \"\"\"\n    Check for temporal biases in structure data.\n    Older structures may have systematic differences from newer ones.\n    \"\"\"\n\n    def __init__(self):\n        self.temporal_stats = None\n\n    def extract_pdb_date(self, target_id: str) -> Optional[datetime]:\n        \"\"\"\n        Try to extract date information from PDB ID.\n        PDB IDs don't directly encode dates, but we can estimate based on ID patterns.\n\n        Note: This is a heuristic. For accurate dates, you'd need to query the PDB.\n        \"\"\"\n        # PDB IDs that start with numbers 1-3 are generally older (1970s-1990s)\n        # IDs starting with 4-9 are more recent (2000s+)\n        # IDs starting with letters are newest (2010s+)\n\n        tid = target_id.upper()\n        if len(tid) >= 4:\n            first_char = tid[0]\n            if first_char.isdigit():\n                num = int(first_char)\n                if num <= 3:\n                    return datetime(1990, 1, 1)  # Old\n                else:\n                    return datetime(2005, 1, 1)  # Medium\n            else:\n                return datetime(2015, 1, 1)  # New\n        return None\n\n    def analyze_temporal_bias(self, seqs_df: pd.DataFrame,\n                              coords_dict: Dict) -> pd.DataFrame:\n        \"\"\"\n        Analyze if there are systematic differences between old and new structures.\n        \"\"\"\n        print(\"\\n=== Temporal Validation ===\")\n\n        temporal_data = []\n\n        for tid, coords in coords_dict.items():\n            if coords is None or len(coords) < 2:\n                continue\n\n            date = self.extract_pdb_date(tid)\n\n            # Compute structure metrics\n            valid_mask = ~(np.isnan(coords).any(axis=1) | np.isinf(coords).any(axis=1))\n            if valid_mask.sum() < 2:\n                continue\n\n            valid_coords = coords[valid_mask]\n            diffs = valid_coords[1:] - valid_coords[:-1]\n            bond_lengths = np.linalg.norm(diffs, axis=1)\n\n            temporal_data.append({\n                'target_id': tid,\n                'estimated_date': date,\n                'era': 'old' if date and date.year < 2000 else ('medium' if date and date.year < 2010 else 'new'),\n                'mean_bond': np.mean(bond_lengths),\n                'std_bond': np.std(bond_lengths),\n                'n_residues': len(valid_coords)\n            })\n\n        self.temporal_stats = pd.DataFrame(temporal_data)\n\n        # Analyze differences between eras\n        print(\"\\n  Bond length statistics by era:\")\n        print(self.temporal_stats.groupby('era')['mean_bond'].agg(['mean', 'std', 'count']))\n\n        return self.temporal_stats\n\n    def filter_by_era(self, coords_dict: Dict,\n                      exclude_eras: List[str] = None) -> Tuple[Dict, List]:\n        \"\"\"\n        Filter structures by estimated era.\n        \"\"\"\n        if self.temporal_stats is None:\n            print(\"Run analyze_temporal_bias first!\")\n            return coords_dict, []\n\n        if exclude_eras is None:\n            return coords_dict, []\n\n        print(f\"\\n=== Filtering by Era (excluding: {exclude_eras}) ===\")\n\n        to_exclude = set(\n            self.temporal_stats[self.temporal_stats['era'].isin(exclude_eras)]['target_id']\n        )\n\n        filtered = {k: v for k, v in coords_dict.items() if k not in to_exclude}\n        removed = list(to_exclude & set(coords_dict.keys()))\n\n        print(f\"  Removed {len(removed)} structures from excluded eras\")\n\n        return filtered, removed","metadata":{"_uuid":"2bf5ad12-15fc-4ba6-a5e8-562aab63bf84","_cell_guid":"26ca4c0a-2b73-4f61-982d-e371bc3104b5","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 12. Cross-validation with PDB","metadata":{"_uuid":"8167b48d-0cef-4e6c-a835-7c3ebda4e05f","_cell_guid":"4c99ba29-9868-433c-96f3-5460774b34fd","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"class PDBValidator:\n    \"\"\"\n    Cross-validate structures against PDB database.\n    \"\"\"\n\n    def __init__(self, pdb_data_path: str = None):\n        \"\"\"\n        Initialize with optional path to local PDB data.\n        \"\"\"\n        self.pdb_data_path = pdb_data_path\n        self.validation_results = {}\n\n    def validate_structure_exists(self, target_id: str) -> bool:\n        \"\"\"\n        Check if structure exists in PDB.\n        This is a placeholder - in practice you'd query the PDB API or local database.\n        \"\"\"\n        # Extract PDB ID (usually first 4 characters)\n        pdb_id = target_id[:4].lower() if len(target_id) >= 4 else target_id.lower()\n\n        # Check local PDB files if path provided\n        if self.pdb_data_path:\n            pdb_file = os.path.join(self.pdb_data_path, f\"{pdb_id}.cif\")\n            return os.path.exists(pdb_file)\n\n        # Otherwise assume valid\n        return True\n\n    def compare_with_pdb(self, target_id: str, coords: np.ndarray,\n                         pdb_coords: np.ndarray = None) -> Dict:\n        \"\"\"\n        Compare structure coordinates with PDB reference.\n        \"\"\"\n        if pdb_coords is None:\n            # In practice, you would load from PDB file\n            return {'rmsd': np.nan, 'valid': True, 'message': 'No PDB reference available'}\n\n        # Align and compute RMSD\n        if len(coords) != len(pdb_coords):\n            return {'rmsd': np.nan, 'valid': False, 'message': 'Length mismatch'}\n\n        # Simple RMSD calculation (without superposition)\n        rmsd = np.sqrt(np.mean(np.sum((coords - pdb_coords)**2, axis=1)))\n\n        return {\n            'rmsd': rmsd,\n            'valid': rmsd < 10.0,  # Threshold for acceptable deviation\n            'message': 'OK' if rmsd < 10.0 else 'Large deviation from PDB'\n        }\n\n    def batch_validate(self, coords_dict: Dict,\n                       seqs_df: pd.DataFrame) -> pd.DataFrame:\n        \"\"\"\n        Validate all structures against PDB.\n        \"\"\"\n        print(\"\\n=== PDB Cross-validation ===\")\n\n        results = []\n        for tid in coords_dict.keys():\n            exists = self.validate_structure_exists(tid)\n            results.append({\n                'target_id': tid,\n                'exists_in_pdb': exists,\n                'pdb_id': tid[:4].lower() if len(tid) >= 4 else tid\n            })\n\n        results_df = pd.DataFrame(results)\n\n        print(f\"  Structures validated: {len(results_df)}\")\n        print(f\"  Confirmed in PDB: {results_df['exists_in_pdb'].sum()}\")\n\n        return results_df","metadata":{"_uuid":"299898b9-7286-4416-95de-b952c2b8cde7","_cell_guid":"c910f176-cda3-4a13-880f-ac58b11f5d29","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 13. Ensemble Averaging","metadata":{"_uuid":"40ed9e51-7ab0-44b9-b452-e07cf2ca314d","_cell_guid":"d7869a6d-92a1-4517-a64b-e57b5a156302","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"class EnsembleAverager:\n    \"\"\"\n    Average coordinates for duplicate/similar structures.\n    \"\"\"\n\n    def __init__(self):\n        self.averaged_structures = {}\n\n    def find_structure_groups(self, seqs_df: pd.DataFrame,\n                              coords_dict: Dict) -> Dict[str, List[str]]:\n        \"\"\"\n        Group structures by sequence identity.\n        \"\"\"\n        seq_to_ids = defaultdict(list)\n\n        for _, row in seqs_df.iterrows():\n            tid = row['target_id']\n            if tid in coords_dict:\n                seq = row['sequence']\n                seq_to_ids[seq].append(tid)\n\n        # Filter to groups with multiple structures\n        groups = {seq: ids for seq, ids in seq_to_ids.items() if len(ids) > 1}\n\n        return groups\n\n    def superpose_structures(self, coords1: np.ndarray,\n                             coords2: np.ndarray) -> Tuple[np.ndarray, float]:\n        \"\"\"\n        Superpose coords2 onto coords1 using Kabsch algorithm.\n        Returns transformed coords2 and RMSD.\n        \"\"\"\n        # Handle NaN\n        valid1 = ~np.isnan(coords1).any(axis=1)\n        valid2 = ~np.isnan(coords2).any(axis=1)\n        valid = valid1 & valid2\n\n        if valid.sum() < 3:\n            return coords2, np.inf\n\n        c1 = coords1[valid]\n        c2 = coords2[valid]\n\n        # Center\n        centroid1 = np.mean(c1, axis=0)\n        centroid2 = np.mean(c2, axis=0)\n        c1_centered = c1 - centroid1\n        c2_centered = c2 - centroid2\n\n        # Compute rotation matrix (Kabsch)\n        H = c2_centered.T @ c1_centered\n        U, S, Vt = np.linalg.svd(H)\n        R = Vt.T @ U.T\n\n        # Handle reflection\n        if np.linalg.det(R) < 0:\n            Vt[-1, :] *= -1\n            R = Vt.T @ U.T\n\n        # Transform coords2\n        coords2_transformed = coords2.copy()\n        coords2_transformed = (coords2_transformed - centroid2) @ R.T + centroid1\n\n        # Compute RMSD\n        rmsd = np.sqrt(np.mean(np.sum((coords1[valid] - coords2_transformed[valid])**2, axis=1)))\n\n        return coords2_transformed, rmsd\n\n    def average_coordinates(self, coord_list: List[np.ndarray]) -> np.ndarray:\n        \"\"\"\n        Average multiple coordinate sets after superposition.\n        \"\"\"\n        if len(coord_list) == 0:\n            return None\n        if len(coord_list) == 1:\n            return coord_list[0]\n\n        # Use first structure as reference\n        reference = coord_list[0]\n        aligned_coords = [reference]\n\n        # Superpose all others onto reference\n        for coords in coord_list[1:]:\n            if len(coords) == len(reference):\n                aligned, rmsd = self.superpose_structures(reference, coords)\n                if rmsd < 20.0:  # Only include if reasonably close\n                    aligned_coords.append(aligned)\n\n        # Average\n        stacked = np.stack(aligned_coords, axis=0)\n        averaged = np.nanmean(stacked, axis=0)\n\n        return averaged\n\n    def process_duplicates(self, seqs_df: pd.DataFrame,\n                           coords_dict: Dict) -> Tuple[Dict, pd.DataFrame]:\n        \"\"\"\n        Process all duplicate groups and create averaged structures.\n        \"\"\"\n        print(\"\\n=== Ensemble Averaging ===\")\n\n        groups = self.find_structure_groups(seqs_df, coords_dict)\n        print(f\"  Found {len(groups)} sequence groups with multiple structures\")\n\n        # Create new coords dict with averaged structures\n        new_coords = coords_dict.copy()\n        averaging_log = []\n\n        for seq, tids in groups.items():\n            coord_list = [coords_dict[tid] for tid in tids if tid in coords_dict]\n\n            if len(coord_list) > 1:\n                averaged = self.average_coordinates(coord_list)\n\n                # Keep first ID, remove others\n                keep_id = tids[0]\n                new_coords[keep_id] = averaged\n\n                for tid in tids[1:]:\n                    if tid in new_coords:\n                        del new_coords[tid]\n\n                averaging_log.append({\n                    'kept_id': keep_id,\n                    'merged_ids': tids[1:],\n                    'n_merged': len(tids)\n                })\n\n        log_df = pd.DataFrame(averaging_log)\n        print(f\"  Averaged {len(averaging_log)} structure groups\")\n        print(f\"  Remaining structures: {len(new_coords)}\")\n\n        return new_coords, log_df","metadata":{"_uuid":"d84fa87d-0545-464b-ac23-8668c5143b4f","_cell_guid":"b25a80c6-04ef-4918-a6f9-4c8aa0899a34","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 14. Chain Break Detection","metadata":{"_uuid":"d549edab-48f9-4b66-bbc2-76cb94c5e114","_cell_guid":"74ff4dcf-5210-456a-be4c-9f41b03a442f","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"class ChainBreakDetector:\n    \"\"\"\n    Detect and handle chain breaks in multi-chain structures.\n    \"\"\"\n\n    def __init__(self, max_bond_distance: float = 9.0):\n        \"\"\"\n        Args:\n            max_bond_distance: Maximum distance between consecutive residues\n                               before considering it a chain break.\n        \"\"\"\n        self.max_bond_distance = max_bond_distance\n        self.chain_info = {}\n\n    def detect_chain_breaks(self, coords: np.ndarray) -> List[int]:\n        \"\"\"\n        Detect positions of chain breaks.\n        Returns list of indices where chain breaks occur (i.e., large gaps).\n        \"\"\"\n        if coords is None or len(coords) < 2:\n            return []\n\n        breaks = []\n        for i in range(len(coords) - 1):\n            if np.isnan(coords[i]).any() or np.isnan(coords[i+1]).any():\n                continue\n\n            dist = np.linalg.norm(coords[i+1] - coords[i])\n            if dist > self.max_bond_distance:\n                breaks.append(i)\n\n        return breaks\n\n    def get_chain_segments(self, coords: np.ndarray) -> List[Tuple[int, int]]:\n        \"\"\"\n        Get list of (start, end) indices for each chain segment.\n        \"\"\"\n        breaks = self.detect_chain_breaks(coords)\n\n        if not breaks:\n            return [(0, len(coords))]\n\n        segments = []\n        prev_end = 0\n\n        for break_idx in breaks:\n            segments.append((prev_end, break_idx + 1))\n            prev_end = break_idx + 1\n\n        segments.append((prev_end, len(coords)))\n\n        return segments\n\n    def analyze_all_structures(self, coords_dict: Dict) -> pd.DataFrame:\n        \"\"\"\n        Analyze chain breaks in all structures.\n        \"\"\"\n        print(\"\\n=== Chain Break Detection ===\")\n\n        results = []\n        for tid, coords in coords_dict.items():\n            breaks = self.detect_chain_breaks(coords)\n            segments = self.get_chain_segments(coords)\n\n            results.append({\n                'target_id': tid,\n                'n_residues': len(coords),\n                'n_chain_breaks': len(breaks),\n                'n_chains': len(segments),\n                'break_positions': breaks,\n                'segment_lengths': [e - s for s, e in segments]\n            })\n\n            self.chain_info[tid] = {\n                'breaks': breaks,\n                'segments': segments\n            }\n\n        results_df = pd.DataFrame(results)\n\n        print(f\"  Analyzed {len(results_df)} structures\")\n        print(f\"  Structures with chain breaks: {(results_df['n_chain_breaks'] > 0).sum()}\")\n        print(f\"  Max chains in a structure: {results_df['n_chains'].max()}\")\n\n        return results_df\n\n    def fix_chain_break_artifacts(self, coords: np.ndarray,\n                                  segments: List[Tuple[int, int]]) -> np.ndarray:\n        \"\"\"\n        Fix artifacts that might occur at chain breaks.\n        Ensures each chain segment is internally consistent.\n        \"\"\"\n        fixed_coords = coords.copy()\n\n        for start, end in segments:\n            if end - start < 2:\n                continue\n\n            # Fix bond lengths within segment\n            segment = fixed_coords[start:end]\n            fixed_segment = fix_bond_lengths(segment)\n            fixed_coords[start:end] = fixed_segment\n\n        return fixed_coords","metadata":{"_uuid":"3b36994a-2cdc-422a-9ebe-961748f841c7","_cell_guid":"7791195d-5869-4ab9-98c5-a2088d893ad6","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 15. Energy Minimization","metadata":{"_uuid":"8376a4cf-58c9-429f-ad1e-f795e4bdb4b8","_cell_guid":"62e84d42-829a-4bbf-91b3-be29546bef55","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"class EnergyMinimizer:\n    \"\"\"\n    Apply simple energy minimization to refine structures.\n    Uses a simplified force field based on bond lengths and angles.\n    \"\"\"\n\n    def __init__(self):\n        self.target_bond = 5.95  # Target C1'-C1' distance\n        self.target_angle = 140.0  # Target bond angle (degrees)\n        self.min_distance = 3.0  # Minimum non-bonded distance\n\n    def compute_energy(self, coords: np.ndarray,\n                       segments: List[Tuple[int, int]] = None) -> Dict:\n        \"\"\"\n        Compute simplified energy terms.\n        \"\"\"\n        if segments is None:\n            segments = [(0, len(coords))]\n\n        energy = {\n            'bond': 0.0,\n            'angle': 0.0,\n            'clash': 0.0,\n            'total': 0.0\n        }\n\n        for start, end in segments:\n            seg = coords[start:end]\n            n = len(seg)\n\n            if n < 2:\n                continue\n\n            # Bond energy\n            diffs = seg[1:] - seg[:-1]\n            bond_lengths = np.linalg.norm(diffs, axis=1)\n            bond_energy = np.sum((bond_lengths - self.target_bond)**2)\n            energy['bond'] += bond_energy\n\n            # Angle energy\n            if n >= 3:\n                v1 = seg[:-2] - seg[1:-1]\n                v2 = seg[2:] - seg[1:-1]\n                cos_angles = np.sum(v1 * v2, axis=1) / (\n                    np.linalg.norm(v1, axis=1) * np.linalg.norm(v2, axis=1) + 1e-8\n                )\n                angles = np.arccos(np.clip(cos_angles, -1, 1)) * 180 / np.pi\n                angle_energy = np.sum((angles - self.target_angle)**2) * 0.01\n                energy['angle'] += angle_energy\n\n        # Clash energy (non-bonded)\n        n_total = len(coords)\n        for i in range(n_total):\n            for j in range(i + 3, n_total):  # Skip bonded atoms\n                if np.isnan(coords[i]).any() or np.isnan(coords[j]).any():\n                    continue\n                dist = np.linalg.norm(coords[i] - coords[j])\n                if dist < self.min_distance:\n                    energy['clash'] += (self.min_distance - dist)**2 * 10\n\n        energy['total'] = energy['bond'] + energy['angle'] + energy['clash']\n\n        return energy\n\n    def minimize(self, coords: np.ndarray,\n                 segments: List[Tuple[int, int]] = None,\n                 max_iterations: int = 100,\n                 learning_rate: float = 0.01,\n                 tolerance: float = 0.1) -> Tuple[np.ndarray, Dict]:\n        \"\"\"\n        Simple gradient descent energy minimization.\n        \"\"\"\n        if segments is None:\n            segments = [(0, len(coords))]\n\n        coords = coords.copy().astype(float)\n        n = len(coords)\n\n        initial_energy = self.compute_energy(coords, segments)\n\n        for iteration in range(max_iterations):\n            # Compute gradient numerically\n            gradient = np.zeros_like(coords)\n            delta = 0.01\n\n            current_energy = self.compute_energy(coords, segments)['total']\n\n            for i in range(n):\n                if np.isnan(coords[i]).any():\n                    continue\n                for j in range(3):\n                    coords[i, j] += delta\n                    e_plus = self.compute_energy(coords, segments)['total']\n                    coords[i, j] -= 2 * delta\n                    e_minus = self.compute_energy(coords, segments)['total']\n                    coords[i, j] += delta\n\n                    gradient[i, j] = (e_plus - e_minus) / (2 * delta)\n\n            # Update coordinates\n            coords -= learning_rate * gradient\n\n            new_energy = self.compute_energy(coords, segments)['total']\n\n            # Check convergence\n            if abs(current_energy - new_energy) < tolerance:\n                break\n\n        final_energy = self.compute_energy(coords, segments)\n\n        return coords, {\n            'initial_energy': initial_energy,\n            'final_energy': final_energy,\n            'iterations': iteration + 1,\n            'converged': iteration < max_iterations - 1\n        }\n\n    def minimize_all(self, coords_dict: Dict,\n                     chain_info: Dict = None,\n                     max_iterations: int = 50) -> Tuple[Dict, pd.DataFrame]:\n        \"\"\"\n        Minimize all structures.\n        \"\"\"\n        print(\"\\n=== Energy Minimization ===\")\n\n        minimized = {}\n        results = []\n\n        for tid, coords in coords_dict.items():\n            segments = chain_info.get(tid, {}).get('segments', [(0, len(coords))]) if chain_info else None\n\n            min_coords, info = self.minimize(coords, segments, max_iterations=max_iterations)\n            minimized[tid] = min_coords\n\n            results.append({\n                'target_id': tid,\n                'initial_energy': info['initial_energy']['total'],\n                'final_energy': info['final_energy']['total'],\n                'energy_reduction': info['initial_energy']['total'] - info['final_energy']['total'],\n                'iterations': info['iterations'],\n                'converged': info['converged']\n            })\n\n        results_df = pd.DataFrame(results)\n\n        print(f\"  Minimized {len(results_df)} structures\")\n        print(f\"  Converged: {results_df['converged'].sum()}\")\n        print(f\"  Mean energy reduction: {results_df['energy_reduction'].mean():.2f}\")\n\n        return minimized, results_df","metadata":{"_uuid":"d5896baf-c837-4ccb-b721-bdda78410f4e","_cell_guid":"5c148bb3-9e64-4e39-b706-2b5d711f10e5","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 16. Complete Data Cleaning Pipeline","metadata":{"_uuid":"59158d5b-5a65-493d-8d59-eeb545f53af5","_cell_guid":"a8754369-8460-4eb8-901d-3e4b6fc1c4c2","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"class ComprehensiveRNACleaner:\n    \"\"\"\n    Complete RNA data cleaning pipeline with all methods.\n    \"\"\"\n\n    def __init__(self, config: Dict = None):\n        self.config = config or {\n            # Basic cleaning\n            'max_nan_fraction': 0.1,\n            'max_coord_range': 2000,\n            'min_bond_length': 3.5,\n            'max_bond_length': 9.0,\n            'target_bond_length': 5.95,\n            'max_bad_bond_fraction': 0.2,\n\n            # Sequence cleaning\n            'min_seq_length': 10,\n            'max_seq_length': 5000,\n            'allowed_nucleotides': set('ACGU'),\n            'max_homopolymer_run': 20,\n            'min_gc_content': 0.15,\n            'max_gc_content': 0.85,\n\n            # Advanced cleaning\n            'resolution_threshold': 2.0,\n            'ss_consistency_threshold': 0.5,\n            'exclude_eras': [],  # e.g., ['old'] to exclude old structures\n            'enable_energy_minimization': True,\n            'enable_ensemble_averaging': True,\n\n            # Output\n            'output_path': OUTPUT_PATH,\n        }\n\n        # Initialize components\n        self.ss_validator = SecondaryStructureValidator()\n        self.resolution_filter = ResolutionFilter()\n        self.temporal_validator = TemporalValidator()\n        self.pdb_validator = PDBValidator()\n        self.ensemble_averager = EnsembleAverager()\n        self.chain_detector = ChainBreakDetector()\n        self.energy_minimizer = EnergyMinimizer()\n\n        # Logging\n        self.cleaning_log = []\n        self.removed_ids = {}\n        self.stats = {}\n\n    def log(self, message: str):\n        self.cleaning_log.append(message)\n        print(message)\n\n    def clean_coordinates_basic(self, coords_dict: Dict) -> Dict:\n        \"\"\"Basic coordinate cleaning.\"\"\"\n        self.log(\"\\n=== Basic Coordinate Cleaning ===\")\n\n        cleaned = {}\n        removed_nan = []\n        removed_range = []\n        removed_bonds = []\n\n        for tid, coords in coords_dict.items():\n            coords = np.array(coords, dtype=float)\n\n            # Check NaN/Inf\n            nan_fraction = np.isnan(coords).sum() / coords.size\n            if nan_fraction > self.config['max_nan_fraction']:\n                removed_nan.append(tid)\n                continue\n\n            # Check coordinate range\n            valid_mask = ~(np.isnan(coords).any(axis=1) | np.isinf(coords).any(axis=1))\n            if valid_mask.sum() > 0:\n                coord_range = np.ptp(coords[valid_mask])\n                if coord_range > self.config['max_coord_range']:\n                    removed_range.append(tid)\n                    continue\n\n            # Check bond lengths\n            if valid_mask.sum() > 1:\n                valid_coords = coords[valid_mask]\n                diffs = valid_coords[1:] - valid_coords[:-1]\n                bond_lengths = np.linalg.norm(diffs, axis=1)\n\n                bad_bonds = np.sum(\n                    (bond_lengths < self.config['min_bond_length']) |\n                    (bond_lengths > self.config['max_bond_length'])\n                )\n                bad_fraction = bad_bonds / len(bond_lengths)\n\n                if bad_fraction > self.config['max_bad_bond_fraction']:\n                    removed_bonds.append(tid)\n                    continue\n\n            cleaned[tid] = coords\n\n        self.removed_ids['nan_coords'] = removed_nan\n        self.removed_ids['coord_range'] = removed_range\n        self.removed_ids['bad_bonds'] = removed_bonds\n\n        self.log(f\"  Removed {len(removed_nan)} structures with too many NaN coordinates\")\n        self.log(f\"  Removed {len(removed_range)} structures with out-of-range coordinates\")\n        self.log(f\"  Removed {len(removed_bonds)} structures with bad bond lengths\")\n        self.log(f\"  Remaining: {len(cleaned)} structures\")\n\n        return cleaned\n\n    def clean_sequences(self, seqs_df: pd.DataFrame) -> pd.DataFrame:\n        \"\"\"Clean sequence data.\"\"\"\n        self.log(\"\\n=== Sequence Cleaning ===\")\n\n        df = seqs_df.copy()\n        initial_count = len(df)\n        removed_reasons = {}\n\n        # Length filter\n        df['seq_length'] = df['sequence'].str.len()\n        length_mask = (\n            (df['seq_length'] >= self.config['min_seq_length']) &\n            (df['seq_length'] <= self.config['max_seq_length'])\n        )\n        removed_reasons['length'] = df[~length_mask]['target_id'].tolist()\n        df = df[length_mask]\n\n        # Non-standard nucleotides\n        def has_only_standard(seq):\n            return set(seq).issubset(self.config['allowed_nucleotides'])\n\n        standard_mask = df['sequence'].apply(has_only_standard)\n        removed_reasons['non_standard'] = df[~standard_mask]['target_id'].tolist()\n        df = df[standard_mask]\n\n        # Homopolymer runs\n        def max_homopolymer(seq):\n            max_run = 1\n            current_run = 1\n            for i in range(1, len(seq)):\n                if seq[i] == seq[i-1]:\n                    current_run += 1\n                    max_run = max(max_run, current_run)\n                else:\n                    current_run = 1\n            return max_run\n\n        df['max_homopolymer'] = df['sequence'].apply(max_homopolymer)\n        homo_mask = df['max_homopolymer'] <= self.config['max_homopolymer_run']\n        removed_reasons['homopolymer'] = df[~homo_mask]['target_id'].tolist()\n        df = df[homo_mask]\n\n        # GC content\n        def gc_content(seq):\n            gc = sum(1 for c in seq if c in 'GC')\n            return gc / len(seq) if len(seq) > 0 else 0\n\n        df['gc_content'] = df['sequence'].apply(gc_content)\n        gc_mask = (\n            (df['gc_content'] >= self.config['min_gc_content']) &\n            (df['gc_content'] <= self.config['max_gc_content'])\n        )\n        removed_reasons['gc_content'] = df[~gc_mask]['target_id'].tolist()\n        df = df[gc_mask]\n\n        self.removed_ids.update({f'seq_{k}': v for k, v in removed_reasons.items()})\n\n        for reason, ids in removed_reasons.items():\n            self.log(f\"  Removed {len(ids)} sequences due to {reason}\")\n\n        self.log(f\"  Remaining: {len(df)} sequences (removed {initial_count - len(df)} total)\")\n\n        return df\n\n    def run_full_pipeline(self, seqs_df: pd.DataFrame,\n                          coords_dict: Dict) -> Tuple[pd.DataFrame, Dict]:\n        \"\"\"\n        Run the complete cleaning pipeline.\n        \"\"\"\n        self.log(\"=\"*70)\n        self.log(\"COMPREHENSIVE RNA DATA CLEANING PIPELINE\")\n        self.log(\"=\"*70)\n        self.log(f\"Initial sequences: {len(seqs_df)}\")\n        self.log(f\"Initial structures: {len(coords_dict)}\")\n\n        # Step 1: Basic sequence cleaning\n        seqs_df = self.clean_sequences(seqs_df)\n\n        # Step 2: Basic coordinate cleaning\n        coords_dict = self.clean_coordinates_basic(coords_dict)\n\n        # Step 3: Chain break detection\n        chain_df = self.chain_detector.analyze_all_structures(coords_dict)\n        self.stats['chain_breaks'] = chain_df\n\n        # Step 4: Resolution-based filtering\n        coords_dict, removed = self.resolution_filter.filter_by_quality(\n            coords_dict, self.config['resolution_threshold']\n        )\n        self.removed_ids['low_resolution'] = removed\n\n        # Step 5: Temporal validation and filtering\n        temporal_df = self.temporal_validator.analyze_temporal_bias(seqs_df, coords_dict)\n        self.stats['temporal'] = temporal_df\n\n        if self.config['exclude_eras']:\n            coords_dict, removed = self.temporal_validator.filter_by_era(\n                coords_dict, self.config['exclude_eras']\n            )\n            self.removed_ids['excluded_eras'] = removed\n\n        # Step 6: Secondary structure validation\n        ss_df = validate_secondary_structures(seqs_df, coords_dict, sample_size=min(200, len(coords_dict)))\n        self.stats['secondary_structure'] = ss_df\n\n        # Filter by SS consistency (optional)\n        if self.config['ss_consistency_threshold'] > 0:\n            low_consistency = ss_df[ss_df['consistency_score'] < self.config['ss_consistency_threshold']]['target_id'].tolist()\n            coords_dict = {k: v for k, v in coords_dict.items() if k not in low_consistency}\n            self.removed_ids['low_ss_consistency'] = low_consistency\n            self.log(f\"  Removed {len(low_consistency)} structures with low SS consistency\")\n\n        # Step 7: Ensemble averaging for duplicates\n        if self.config['enable_ensemble_averaging']:\n            coords_dict, avg_log = self.ensemble_averager.process_duplicates(seqs_df, coords_dict)\n            self.stats['ensemble_averaging'] = avg_log\n\n        # Step 8: Energy minimization (DISABLED BY DEFAULT - very slow)\n        # This uses gradient descent which is O(n^2) per iteration\n        # For large datasets, consider using faster optimization methods\n        if self.config['enable_energy_minimization']:\n            self.log(\"\\n=== Energy Minimization (SLOW - Consider disabling) ===\")\n            coords_dict, min_log = self.energy_minimizer.minimize_all(\n                coords_dict, self.chain_detector.chain_info, max_iterations=30\n            )\n            self.stats['minimization'] = min_log\n        else:\n            self.log(\"\\n=== Energy Minimization: SKIPPED (disabled in config) ===\")\n\n        # Step 9: Fix bond lengths\n        self.log(\"\\n=== Final Bond Length Correction ===\")\n        fixed_count = 0\n        for tid in coords_dict:\n            segments = self.chain_detector.chain_info.get(tid, {}).get('segments', [(0, len(coords_dict[tid]))])\n            coords_dict[tid] = self.chain_detector.fix_chain_break_artifacts(coords_dict[tid], segments)\n            fixed_count += 1\n        self.log(f\"  Fixed bond lengths for {fixed_count} structures\")\n\n        # Step 10: Align sequences and coordinates\n        self.log(\"\\n=== Final Alignment ===\")\n        common_ids = set(seqs_df['target_id']) & set(coords_dict.keys())\n        seqs_df = seqs_df[seqs_df['target_id'].isin(common_ids)].copy()\n        coords_dict = {k: v for k, v in coords_dict.items() if k in common_ids}\n\n        self.log(\"\\n\" + \"=\"*70)\n        self.log(\"CLEANING COMPLETE\")\n        self.log(\"=\"*70)\n        self.log(f\"Final sequences: {len(seqs_df)}\")\n        self.log(f\"Final structures: {len(coords_dict)}\")\n\n        return seqs_df, coords_dict\n\n    def get_cleaning_summary(self) -> pd.DataFrame:\n        \"\"\"Get summary of all cleaning operations.\"\"\"\n        summary = pd.DataFrame([\n            {'Reason': k, 'Count': len(v) if isinstance(v, list) else 0}\n            for k, v in self.removed_ids.items()\n        ])\n        return summary\n\n    def _coords_to_pdb_string(self, target_id: str, sequence: str,\n                               coords: np.ndarray) -> str:\n        \"\"\"\n        Convert cleaned coordinates to a PDB-format string.\n\n        Each residue is represented by its C1' atom (the coordinate we have).\n        This produces valid PDB files that SSR and RNApdbee can parse.\n        \"\"\"\n        # Map single-letter to 3-letter residue names\n        res_map = {'A': '  A', 'C': '  C', 'G': '  G', 'U': '  U'}\n\n        lines = []\n        lines.append(f\"HEADER    RNA STRUCTURE                             {target_id}\")\n        atom_serial = 1\n        chain_id = 'A'\n\n        for i, (x, y, z) in enumerate(coords):\n            if np.isnan(x) or np.isnan(y) or np.isnan(z):\n                continue\n            resname = res_map.get(sequence[i] if i < len(sequence) else 'A', '  A')\n            resid = i + 1\n            # PDB ATOM record format (fixed-width columns)\n            line = (\n                f\"ATOM  {atom_serial:5d}  C1'{resname} {chain_id}{resid:4d}    \"\n                f\"{x:8.3f}{y:8.3f}{z:8.3f}  1.00  0.00           C\"\n            )\n            lines.append(line)\n            atom_serial += 1\n\n        lines.append(\"END\")\n        return '\\n'.join(lines) + '\\n'\n\n    def export_cleaned_pdbs(self, seqs_df: pd.DataFrame, coords_dict: Dict,\n                            output_dir: str = None) -> str:\n        \"\"\"\n        Export each cleaned structure as an individual PDB file.\n        Returns the output directory path.\n        \"\"\"\n        if output_dir is None:\n            output_dir = os.path.join(self.config['output_path'], 'cleaned_pdbs')\n        os.makedirs(output_dir, exist_ok=True)\n\n        self.log(f\"\\n=== Exporting Cleaned PDB Files to {output_dir} ===\")\n        count = 0\n        for tid, coords in coords_dict.items():\n            seq_row = seqs_df[seqs_df['target_id'] == tid]\n            if len(seq_row) == 0:\n                continue\n            sequence = seq_row.iloc[0]['sequence']\n            pdb_str = self._coords_to_pdb_string(tid, sequence, coords)\n            pdb_file = os.path.join(output_dir, f\"{tid}.pdb\")\n            with open(pdb_file, 'w') as f:\n                f.write(pdb_str)\n            count += 1\n\n        self.log(f\"  Exported {count} PDB files\")\n        return output_dir\n\n    def save_cleaned_data(self, seqs_df: pd.DataFrame, coords_dict: Dict,\n                          output_prefix: str = 'cleaned'):\n        \"\"\"\n        Save cleaned data to files.\n        \"\"\"\n        output_path = self.config['output_path']\n\n        self.log(\"\\n\" + \"=\"*70)\n        self.log(\"SAVING CLEANED DATA\")\n        self.log(\"=\"*70)\n\n        # Save sequences\n        seq_file = os.path.join(output_path, f'{output_prefix}_sequences.csv')\n        seqs_df.to_csv(seq_file, index=False)\n        self.log(f\"  Saved sequences to: {seq_file}\")\n\n        # Reconstruct and save labels\n        labels_list = []\n        for tid, coords in coords_dict.items():\n            seq_row = seqs_df[seqs_df['target_id'] == tid]\n            if len(seq_row) == 0:\n                continue\n            seq = seq_row.iloc[0]['sequence']\n\n            for i, (x, y, z) in enumerate(coords):\n                resname = seq[i] if i < len(seq) else 'X'\n                labels_list.append({\n                    'ID': f\"{tid}_{i+1}\",\n                    'resname': resname,\n                    'resid': i + 1,\n                    'x_1': x,\n                    'y_1': y,\n                    'z_1': z\n                })\n\n        labels_df = pd.DataFrame(labels_list)\n        labels_file = os.path.join(output_path, f'{output_prefix}_labels.csv')\n        labels_df.to_csv(labels_file, index=False)\n        self.log(f\"  Saved labels to: {labels_file}\")\n\n        # Save cleaning summary\n        summary = self.get_cleaning_summary()\n        summary_file = os.path.join(output_path, f'{output_prefix}_cleaning_summary.csv')\n        summary.to_csv(summary_file, index=False)\n        self.log(f\"  Saved cleaning summary to: {summary_file}\")\n\n        # Save cleaning log\n        log_file = os.path.join(output_path, f'{output_prefix}_cleaning_log.txt')\n        with open(log_file, 'w') as f:\n            f.write('\\n'.join(self.cleaning_log))\n        self.log(f\"  Saved cleaning log to: {log_file}\")\n\n        # Save statistics\n        stats_file = os.path.join(output_path, f'{output_prefix}_stats.json')\n        stats_to_save = {}\n        for key, val in self.stats.items():\n            if isinstance(val, pd.DataFrame):\n                stats_to_save[key] = val.to_dict()\n        with open(stats_file, 'w') as f:\n            json.dump(stats_to_save, f, indent=2, default=str)\n        self.log(f\"  Saved statistics to: {stats_file}\")\n\n        # Export individual PDB files for downstream SS extraction\n        pdb_dir = self.export_cleaned_pdbs(seqs_df, coords_dict)\n        saved_files_extra = {'pdb_directory': pdb_dir}\n\n        self.log(\"\\n  All files saved successfully!\")\n\n        return {\n            'sequences_file': seq_file,\n            'labels_file': labels_file,\n            'summary_file': summary_file,\n            'log_file': log_file,\n            'stats_file': stats_file,\n            'pdb_directory': pdb_dir,\n        }","metadata":{"_uuid":"b6a8aae4-46f3-4551-8c50-d73556fc71d1","_cell_guid":"65a54071-5a30-4203-8d07-51e960f54d8f","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## Pre-Pipeline Analysis: Understand What Would Be Filtered\n#\nBefore running the pipeline, let's see what each filter would remove","metadata":{"_uuid":"73d18fb0-196c-4096-a105-a106db93f4fa","_cell_guid":"8135d39f-e4b6-444d-b7d1-6982c86832e9","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"# EXPLORATORY ANALYSIS - See what would be filtered\nprint(\"\\n\" + \"=\"*70)\nprint(\"PRE-PIPELINE ANALYSIS: Understanding Filter Impact\")\nprint(\"=\"*70)\n\n# Sequence length distribution\nprint(\"\\n--- Sequence Length Analysis ---\")\nseq_lengths = train_seqs['sequence'].str.len()\nprint(f\"Min length: {seq_lengths.min()}\")\nprint(f\"Max length: {seq_lengths.max()}\")\nprint(f\"Mean length: {seq_lengths.mean():.1f}\")\nprint(f\"Median length: {seq_lengths.median():.1f}\")\nprint(f\"Sequences < 10 nt: {(seq_lengths < 10).sum()}\")\nprint(f\"Sequences > 5000 nt: {(seq_lengths > 5000).sum()}\")\n\n# Homopolymer analysis\nprint(\"\\n--- Homopolymer Run Analysis ---\")\ndef get_max_homopolymer(seq):\n    if len(seq) < 2:\n        return 1\n    max_run = 1\n    current_run = 1\n    for i in range(1, len(seq)):\n        if seq[i] == seq[i-1]:\n            current_run += 1\n            max_run = max(max_run, current_run)\n        else:\n            current_run = 1\n    return max_run\n\nhomopolymer_runs = train_seqs['sequence'].apply(get_max_homopolymer)\nprint(f\"Max homopolymer run in dataset: {homopolymer_runs.max()}\")\nprint(f\"Mean homopolymer run: {homopolymer_runs.mean():.1f}\")\nprint(f\"Sequences with runs > 10: {(homopolymer_runs > 10).sum()}\")\nprint(f\"Sequences with runs > 15: {(homopolymer_runs > 15).sum()}\")\nprint(f\"Sequences with runs > 20: {(homopolymer_runs > 20).sum()}\")\n\n# GC content analysis\nprint(\"\\n--- GC Content Analysis ---\")\ndef get_gc_content(seq):\n    gc = sum(1 for c in seq if c in 'GC')\n    return gc / len(seq) if len(seq) > 0 else 0\n\ngc_contents = train_seqs['sequence'].apply(get_gc_content)\nprint(f\"Min GC content: {gc_contents.min():.3f}\")\nprint(f\"Max GC content: {gc_contents.max():.3f}\")\nprint(f\"Mean GC content: {gc_contents.mean():.3f}\")\nprint(f\"Sequences with GC < 15%: {(gc_contents < 0.15).sum()}\")\nprint(f\"Sequences with GC > 85%: {(gc_contents > 0.85).sum()}\")\n\n# NaN coordinates analysis\nprint(\"\\n--- Coordinate Quality Analysis ---\")\nnan_counts = []\nfor tid, coords in train_coords_dict.items():\n    nan_frac = np.isnan(coords).sum() / coords.size\n    nan_counts.append({'target_id': tid, 'nan_fraction': nan_frac})\nnan_df = pd.DataFrame(nan_counts)\nprint(f\"Structures with >10% NaN: {(nan_df['nan_fraction'] > 0.1).sum()}\")\nprint(f\"Structures with >50% NaN: {(nan_df['nan_fraction'] > 0.5).sum()}\")\nprint(f\"Structures with >90% NaN: {(nan_df['nan_fraction'] > 0.9).sum()}\")\n\nprint(\"\\n\" + \"=\"*70)\nprint(\"RECOMMENDATION: The main issue is NaN coordinates (1033 structures)\")\nprint(\"Other filters (length, homopolymer, GC) are now disabled to preserve data\")\nprint(\"=\"*70)","metadata":{"_uuid":"ac6c522c-3d19-44e7-a18f-75d780184c78","_cell_guid":"36c82cf7-6b1e-44a3-9ec8-69c4940aa0ae","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## Run the Complete Pipeline","metadata":{"_uuid":"9e57f9a6-721c-4ed7-98fe-aaf1876598d5","_cell_guid":"36b3cbab-017a-436f-8484-b27934533463","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"# Initialize cleaner with configuration\n# NOTE: Relaxed filters to avoid losing too many valid structures\ncleaner = ComprehensiveRNACleaner({\n    # Basic cleaning\n    'max_nan_fraction': 0.1,\n    'max_coord_range': 2000,\n    'min_bond_length': 3.5,\n    'max_bond_length': 9.0,\n    'target_bond_length': 5.95,\n    'max_bad_bond_fraction': 0.2,\n\n    # Sequence cleaning - RELAXED FILTERS\n    # We keep ALL valid RNA sequences - no arbitrary length/composition filters\n    'min_seq_length': 1,          # Keep all lengths - short RNAs are valid\n    'max_seq_length': 100000,     # No upper limit\n    'allowed_nucleotides': set('ACGU'),  # Only filter non-standard nucleotides\n    'max_homopolymer_run': 1000,  # Disabled - homopolymers are valid in RNA\n    'min_gc_content': 0.0,        # Disabled - extreme GC content is valid\n    'max_gc_content': 1.0,        # Disabled - extreme GC content is valid\n\n    # Advanced cleaning - RELAXED THRESHOLDS\n    'resolution_threshold': 10.0,  # Much more lenient (was 2.0, removed 43%!)\n    'ss_consistency_threshold': 0.0,  # DISABLED - heuristic SS is unreliable\n    'exclude_eras': [],  # Don't exclude based on age\n    'enable_energy_minimization': False,  # DISABLED - too slow\n    'enable_ensemble_averaging': True,  # Keep - helps with duplicates\n\n    # Output\n    'output_path': OUTPUT_PATH,\n})\n\n# Run the pipeline\ncleaned_seqs, cleaned_coords = cleaner.run_full_pipeline(\n    train_seqs.copy(),\n    train_coords_dict.copy()\n)","metadata":{"_uuid":"ce51500e-c8e6-4a4f-b029-483e1dc8ef88","_cell_guid":"b3b5064b-3a57-4dde-9923-f792ea55e495","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# View cleaning summary\nprint(\"\\n\" + \"=\"*70)\nprint(\"CLEANING SUMMARY\")\nprint(\"=\"*70)\nprint(cleaner.get_cleaning_summary())","metadata":{"_uuid":"db4d9885-a413-4eb9-a9c6-d2720b599710","_cell_guid":"ddc17db3-6575-47a9-bae5-c75f54315cef","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# SAVE THE CLEANED DATA\nsaved_files = cleaner.save_cleaned_data(\n    cleaned_seqs,\n    cleaned_coords,\n    output_prefix='cleaned'\n)\n\nprint(\"\\n\" + \"=\"*70)\nprint(\"OUTPUT FILES\")\nprint(\"=\"*70)\nfor name, path in saved_files.items():\n    print(f\"  {name}: {path}\")","metadata":{"_uuid":"c515eb7f-9326-47be-827b-37a437a7d2f8","_cell_guid":"d61d37d8-6eba-42b7-8ca3-9f65103c9d61","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## Summary of All Cleaning Methods\n#\n### Basic Cleaning:\n1. **Missing Value Handling** - NaN/Inf detection and removal\n2. **Coordinate Quality** - Range validation, bond length validation\n3. **Sequence Quality** - Length, nucleotide composition, GC content, homopolymers\n4. **Outlier Detection** - IQR and Z-score methods\n5. **Duplicate Removal** - Exact and similarity-based\n#\n### Advanced Cleaning:\n6. **Secondary Structure Validation** - RNAfold prediction vs 3D structure\n7. **Resolution-based Filtering** - Quality metrics proxy\n8. **Temporal Validation** - Era-based bias detection\n9. **PDB Cross-validation** - Validate against original entries\n10. **Ensemble Averaging** - Average duplicate structures\n11. **Chain Break Detection** - Multi-chain handling\n12. **Energy Minimization** - Molecular mechanics refinement\n#\n### Output Files:\n- `cleaned_sequences.csv` - Cleaned sequence data\n- `cleaned_labels.csv` - Cleaned coordinate labels\n- `cleaned_cleaning_summary.csv` - Summary of removed structures\n- `cleaned_cleaning_log.txt` - Detailed cleaning log\n- `cleaned_stats.json` - Statistics from each cleaning step","metadata":{"_uuid":"ce7dd4d9-912d-47b2-be36-8069f90d051f","_cell_guid":"43c9a864-e70a-415d-a6c0-e99eac0d3255","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"# Print removed ids\nprint(\"\\n\" + \"=\"*70)\nprint(\"REMOVED STRUCTURES AND SEQUENCES IDS\")\nprint(\"=\"*70)\n\nfor reason, ids in cleaner.removed_ids.items():\n    if ids:\n        print(f\"\\n{reason.upper()}:\")\n        for id in ids:\n            print(f\"  {id}\")","metadata":{"_uuid":"97ea4bf5-254d-4aeb-a213-5a087150ef00","_cell_guid":"421a847a-323b-4c73-8432-7e148376948b","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"cc9de5b4-0786-4ec8-a009-e840d4ea958d","_cell_guid":"d5a62786-5fb4-4ec1-93bb-2176a8ce817b","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null}]}