{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":87793,"databundleVersionId":11553390,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":11475983,"sourceType":"datasetVersion","datasetId":7192371}],"dockerImageVersionId":31011,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"#!pip install Bio","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:01:10.505942Z","iopub.execute_input":"2025-04-19T16:01:10.506181Z","iopub.status.idle":"2025-04-19T16:01:17.758291Z","shell.execute_reply.started":"2025-04-19T16:01:10.506157Z","shell.execute_reply":"2025-04-19T16:01:17.757574Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile model.py\n\nimport numpy as np\nimport pandas as pd\nimport os\nimport re\nfrom Bio import SeqIO\nfrom Bio.Seq import Seq\nfrom scipy.spatial.transform import Rotation\nfrom scipy.spatial import distance\nfrom scipy.optimize import minimize\n\n# Constants for RNA structure\nNUCLEOTIDE_PAIRS = {\n    'A': 'U',\n    'U': 'A',\n    'G': 'C',\n    'C': 'G'\n}\n\n# Distance between adjacent nucleotides (C1' atoms) in Angstroms\nADJACENT_NUCLEOTIDE_DISTANCE = 6.0  \n\n# Distance between paired nucleotides (C1' atoms) in Angstroms \nBASE_PAIR_DISTANCE = 18.0\n\n# Enhanced structural variation with RNA-specific constraints\ndef enhanced_structural_variation(base_structure, noise_level=0.5, preserve_distance=True, \n                               use_global_movement=False, correlation=0.8, sequence=None):\n    \"\"\"\n    Generate structural variations with RNA-specific constraints.\n    \n    Parameters:\n    - base_structure: Base coordinates to vary\n    - noise_level: Level of random noise (higher = more variation)\n    - preserve_distance: Maintain distances between adjacent residues\n    - use_global_movement: Apply global movements (useful for larger RNAs)\n    - correlation: Correlation between movements of adjacent residues\n    - sequence: RNA sequence for sequence-aware variations\n    \n    Returns:\n    - Varied structure coordinates\n    \"\"\"\n    if base_structure is None or len(base_structure) == 0:\n        return base_structure\n    \n    # Copy the base structure\n    result = base_structure.copy()\n    seq_length = len(result)\n    \n    # RNA-specific local structural patterns\n    if use_global_movement:\n        # Apply bending or twisting to simulate larger conformational changes\n        # Get a random rotation axis\n        rotation_axis = np.random.normal(0, 1, 3)\n        rotation_axis = rotation_axis / np.linalg.norm(rotation_axis)\n        \n        # Center the structure\n        center = np.mean(result, axis=0)\n        centered = result - center\n        \n        for i in range(seq_length):\n            # Calculate gradual rotation based on position\n            rotation_factor = i / seq_length * np.pi * noise_level\n            \n            # Create rotation matrix\n            rot = Rotation.from_rotvec(rotation_axis * rotation_factor)\n            rotation_matrix = rot.as_matrix()\n            \n            # Apply rotation\n            result[i] = np.dot(centered[i], rotation_matrix) + center\n            \n            # Add small noise based on position\n            position_noise = np.random.normal(0, noise_level * 0.5, 3)\n            result[i] += position_noise\n    \n    # Apply sequence-specific noise\n    if sequence is not None:\n        for i in range(seq_length):\n            if i < len(sequence):\n                # Different nucleotides have different flexibility patterns\n                if sequence[i] == 'A' or sequence[i] == 'U':\n                    # A and U tend to be more flexible\n                    local_noise = noise_level * 1.2\n                elif sequence[i] == 'G' or sequence[i] == 'C':\n                    # G and C tend to be more rigid due to stronger base pairing\n                    local_noise = noise_level * 0.8\n                else:\n                    local_noise = noise_level\n                \n                # Add nucleotide-specific noise\n                result[i] += np.random.normal(0, local_noise, 3)\n    \n    # Add correlated noise to simulate connected movement\n    correlated_noise = np.zeros_like(result)\n    \n    # Generate base noise\n    base_noise = np.random.normal(0, noise_level, (seq_length, 3))\n    \n    for i in range(seq_length):\n        if i == 0:\n            correlated_noise[i] = base_noise[i]\n        else:\n            # Correlation with previous residue\n            correlated_noise[i] = correlation * correlated_noise[i-1] + (1-correlation) * base_noise[i]\n    \n    # Apply the correlated noise\n    result += correlated_noise\n    \n    # Preserve distances between adjacent residues if required\n    if preserve_distance and seq_length > 1:\n        for i in range(seq_length - 1):\n            # Get the current vector between adjacent residues\n            vec = result[i+1] - result[i]\n            current_dist = np.linalg.norm(vec)\n            \n            # Target distance should be the normal C1'-C1' distance\n            if current_dist > 0:  # Avoid division by zero\n                # Scale the vector to maintain proper distance\n                vec_scaled = vec * (ADJACENT_NUCLEOTIDE_DISTANCE / current_dist)\n                \n                # Update the position of the next residue\n                result[i+1] = result[i] + vec_scaled\n    \n    return result\n\ndef apply_base_pairing_constraints(coords, sequence, min_loop_size=3):\n    \"\"\"\n    Apply base pairing constraints to the structure by predicting\n    potential base pairs from the sequence and adjusting coordinates.\n    \n    Parameters:\n    - coords: 3D coordinates array\n    - sequence: RNA sequence\n    - min_loop_size: Minimum number of nucleotides in a loop\n    \n    Returns:\n    - Adjusted coordinates that better respect RNA base pairing\n    \"\"\"\n    if len(coords) != len(sequence):\n        return coords  # Cannot apply constraints if lengths don't match\n    \n    result = coords.copy()\n    n = len(sequence)\n    \n    # Simple base pair prediction using complementary bases\n    pairs = []\n    for i in range(n):\n        for j in range(i + min_loop_size + 1, n):\n            if are_complementary(sequence[i], sequence[j]):\n                # The further apart, the more likely a base pair (weighted by GC content)\n                if sequence[i] in 'GC':\n                    weight = 1.2  # GC pairs are stronger\n                else:\n                    weight = 0.8\n                    \n                pairs.append((i, j, weight))\n    \n    # Sort pairs by weight (stronger pairs first)\n    pairs.sort(key=lambda x: x[2], reverse=True)\n    \n    # Track which residues are already paired\n    paired = set()\n    final_pairs = []\n    \n    # Greedy algorithm to select base pairs\n    for i, j, weight in pairs:\n        if i not in paired and j not in paired:\n            paired.add(i)\n            paired.add(j)\n            final_pairs.append((i, j))\n    \n    # Adjust coordinates based on selected base pairs\n    for i, j in final_pairs:\n        # Current distance\n        dist = np.linalg.norm(result[i] - result[j])\n        \n        # If the distance isn't close to the expected base pair distance\n        if abs(dist - BASE_PAIR_DISTANCE) > 3.0:\n            # Calculate midpoint\n            midpoint = (result[i] + result[j]) / 2\n            \n            # Get direction vector\n            direction = result[j] - result[i]\n            if np.linalg.norm(direction) > 0:\n                direction = direction / np.linalg.norm(direction)\n            else:\n                # If residues are at the same position, use a random direction\n                direction = np.random.normal(0, 1, 3)\n                direction = direction / np.linalg.norm(direction)\n            \n            # Move residues to be BASE_PAIR_DISTANCE apart\n            result[i] = midpoint - direction * (BASE_PAIR_DISTANCE / 2)\n            result[j] = midpoint + direction * (BASE_PAIR_DISTANCE / 2)\n    \n    return result\n\ndef are_complementary(n1, n2):\n    \"\"\"Check if nucleotides form a valid base pair\"\"\"\n    return (n1 == 'A' and n2 == 'U') or \\\n           (n1 == 'U' and n2 == 'A') or \\\n           (n1 == 'G' and n2 == 'C') or \\\n           (n1 == 'C' and n2 == 'G')\n\ndef analyze_msa_for_covariation(msa_file):\n    \"\"\"\n    Analyze a Multiple Sequence Alignment file to identify covarying positions\n    which likely represent base pairs.\n    \n    Parameters:\n    - msa_file: Path to MSA file in FASTA format\n    \n    Returns:\n    - List of tuples (i,j) indicating positions that likely form base pairs\n    \"\"\"\n    try:\n        # Parse MSA file\n        sequences = []\n        with open(msa_file, 'r') as f:\n            for record in SeqIO.parse(f, 'fasta'):\n                sequences.append(str(record.seq).upper())\n        \n        if not sequences:\n            return []\n        \n        n = len(sequences[0])\n        covariation_matrix = np.zeros((n, n))\n        \n        # Simple mutual information calculation for covariation\n        for i in range(n):\n            for j in range(i + 3, n):  # Min loop size of 3\n                # Count nucleotide frequencies\n                i_counts = {'A': 0, 'C': 0, 'G': 0, 'U': 0, 'T': 0, '-': 0, 'N': 0}\n                j_counts = {'A': 0, 'C': 0, 'G': 0, 'U': 0, 'T': 0, '-': 0, 'N': 0}\n                pair_counts = {}\n                \n                valid_seqs = 0\n                for seq in sequences:\n                    if i < len(seq) and j < len(seq) and seq[i] != '-' and seq[j] != '-' and seq[i] != 'N' and seq[j] != 'N':\n                        valid_seqs += 1\n                        \n                        # Convert T to U for RNA\n                        i_nuc = 'U' if seq[i] == 'T' else seq[i]\n                        j_nuc = 'U' if seq[j] == 'T' else seq[j]\n                        \n                        i_counts[i_nuc] += 1\n                        j_counts[j_nuc] += 1\n                        \n                        pair = (i_nuc, j_nuc)\n                        pair_counts[pair] = pair_counts.get(pair, 0) + 1\n                \n                # Skip if not enough valid sequences\n                if valid_seqs < 5:\n                    continue\n                \n                # Check for complementary base pair conservation\n                complementary_pairs = 0\n                for pair, count in pair_counts.items():\n                    if are_complementary(pair[0], pair[1]):\n                        complementary_pairs += count\n                \n                covariation_matrix[i, j] = complementary_pairs / valid_seqs\n                covariation_matrix[j, i] = covariation_matrix[i, j]\n        \n        # Extract likely base pairs\n        likely_pairs = []\n        for i in range(n):\n            for j in range(i + 3, n):\n                if covariation_matrix[i, j] > 0.6:  # Threshold for base pair prediction\n                    likely_pairs.append((i, j))\n        \n        return likely_pairs\n    \n    except Exception as e:\n        print(f\"Error analyzing MSA file: {e}\")\n        return []\n\ndef apply_msa_constraints(coords, sequence, msa_file):\n    \"\"\"\n    Apply constraints from MSA analysis to the structure\n    \n    Parameters:\n    - coords: 3D coordinates array\n    - sequence: RNA sequence\n    - msa_file: Path to MSA file\n    \n    Returns:\n    - Adjusted coordinates based on MSA analysis\n    \"\"\"\n    if msa_file is None or not os.path.exists(msa_file):\n        return coords\n    \n    result = coords.copy()\n    \n    try:\n        # Get likely base pairs from MSA\n        likely_pairs = analyze_msa_for_covariation(msa_file)\n        \n        # Apply base pair constraints\n        for i, j in likely_pairs:\n            if i < len(coords) and j < len(coords):\n                # Current distance\n                dist = np.linalg.norm(result[i] - result[j])\n                \n                # If the distance isn't close to the expected base pair distance\n                if abs(dist - BASE_PAIR_DISTANCE) > 3.0:\n                    # Calculate midpoint\n                    midpoint = (result[i] + result[j]) / 2\n                    \n                    # Get direction vector\n                    direction = result[j] - result[i]\n                    if np.linalg.norm(direction) > 0:\n                        direction = direction / np.linalg.norm(direction)\n                    else:\n                        # If residues are at the same position, use a random direction\n                        direction = np.random.normal(0, 1, 3)\n                        direction = direction / np.linalg.norm(direction)\n                    \n                    # Move residues to be BASE_PAIR_DISTANCE apart\n                    result[i] = midpoint - direction * (BASE_PAIR_DISTANCE / 2)\n                    result[j] = midpoint + direction * (BASE_PAIR_DISTANCE / 2)\n    \n    except Exception as e:\n        print(f\"Error applying MSA constraints: {e}\")\n    \n    return result\n\ndef normalize_structure(structure):\n    \"\"\"\n    Normalizes the structure by centering and scaling.\n    \n    Parameters:\n    - structure: 3D coordinates array\n    \n    Returns:\n    - Normalized structure\n    \"\"\"\n    if structure is None or len(structure) == 0:\n        return structure\n    \n    result = structure.copy()\n    \n    # Center the structure\n    center = np.mean(result, axis=0)\n    result = result - center\n    \n    # Scale to a consistent size\n    max_dist = 0\n    for coord in result:\n        dist = np.linalg.norm(coord)\n        if dist > max_dist:\n            max_dist = dist\n    \n    if max_dist > 0:\n        # Scale the structure to a consistent radius\n        result = result * (50.0 / max_dist)\n    \n    return result\n\ndef generate_diverse_ensemble(base_structure, sequence, msa_file=None, num_structures=5):\n    \"\"\"\n    Generate a diverse ensemble of structures for a given RNA sequence\n    \n    Parameters:\n    - base_structure: Base coordinates to start from\n    - sequence: RNA sequence\n    - msa_file: Optional path to MSA file for that sequence\n    - num_structures: Number of structures to generate (default: 5)\n    \n    Returns:\n    - List of structures (each with 3D coordinates)\n    \"\"\"\n    structures = []\n    seq_length = len(sequence)\n    \n    # First structure: Apply base pairing and MSA constraints\n    base_with_constraints = apply_base_pairing_constraints(base_structure, sequence)\n    if msa_file:\n        base_with_constraints = apply_msa_constraints(base_with_constraints, sequence, msa_file)\n    structures.append(normalize_structure(base_with_constraints))\n    \n    # Second structure: Higher noise with global movement\n    structure2 = enhanced_structural_variation(\n        base_structure, \n        noise_level=0.8, \n        preserve_distance=True,\n        use_global_movement=True,\n        sequence=sequence\n    )\n    structure2 = apply_base_pairing_constraints(structure2, sequence)\n    structures.append(normalize_structure(structure2))\n    \n    # Third structure: Medium noise, no global movement\n    structure3 = enhanced_structural_variation(\n        base_structure, \n        noise_level=0.5, \n        preserve_distance=True,\n        use_global_movement=False,\n        sequence=sequence\n    )\n    structure3 = apply_base_pairing_constraints(structure3, sequence)\n    structures.append(normalize_structure(structure3))\n    \n    # Fourth structure: Low noise but different rotation\n    structure4 = base_structure.copy()\n    \n    # Apply random rotation\n    center = np.mean(structure4, axis=0)\n    centered = structure4 - center\n    \n    # Random rotation matrix\n    rotation = Rotation.random()\n    rotation_matrix = rotation.as_matrix()\n    \n    # Apply rotation and re-center\n    for i in range(len(structure4)):\n        structure4[i] = np.dot(centered[i], rotation_matrix) + center\n    \n    structure4 = enhanced_structural_variation(\n        structure4, \n        noise_level=0.3, \n        preserve_distance=True,\n        use_global_movement=False,\n        sequence=sequence\n    )\n    structure4 = apply_base_pairing_constraints(structure4, sequence)\n    structures.append(normalize_structure(structure4))\n    \n    # Fifth structure: Significantly different conformation\n    # Create a different starting point with a more open conformation\n    structure5 = np.zeros_like(base_structure)\n    \n    # Create an elongated starting structure\n    for i in range(seq_length):\n        structure5[i] = np.array([i * ADJACENT_NUCLEOTIDE_DISTANCE * 0.8, 0, 0])\n    \n    # Apply high noise and global movement\n    structure5 = enhanced_structural_variation(\n        structure5, \n        noise_level=1.0, \n        preserve_distance=True,\n        use_global_movement=True,\n        sequence=sequence\n    )\n    structure5 = apply_base_pairing_constraints(structure5, sequence)\n    structures.append(normalize_structure(structure5))\n    \n    # Ensure we have exactly 5 structures\n    while len(structures) < 5:\n        # Create additional structures if needed\n        noise = 0.5 + len(structures) * 0.2\n        new_structure = enhanced_structural_variation(\n            base_structure, \n            noise_level=noise, \n            preserve_distance=True,\n            use_global_movement=(len(structures) % 2 == 0),\n            sequence=sequence\n        )\n        new_structure = apply_base_pairing_constraints(new_structure, sequence)\n        structures.append(normalize_structure(new_structure))\n    \n    return structures[:5]  # Return exactly 5 structures\n\nclass ImprovedRNAModel:\n    \"\"\"\n    Improved RNA 3D structure prediction model that incorporates:\n    - Reference-based modeling\n    - Sequence-specific structural features\n    - MSA-based constraints\n    - RNA-specific geometry constraints\n    - Ensemble prediction\n    \"\"\"\n    \n    def __init__(self, geometric_sampling=True, base_noise_level=0.5, correlation=0.8, msa_dir=None):\n        \"\"\"\n        Initialize the RNA 3D structure model\n        \n        Parameters:\n        - geometric_sampling: Whether to use geometric sampling for structure variation\n        - base_noise_level: Base level of noise for variations\n        - correlation: Correlation between adjacent residues' movements\n        - msa_dir: Directory containing MSA files\n        \"\"\"\n        self.geometric_sampling = geometric_sampling\n        self.base_noise_level = base_noise_level\n        self.correlation = correlation\n        self.msa_dir = msa_dir\n        \n        # Will be populated during fit\n        self.reference_structures = []\n        self.reference_sequences = []\n        self.size_groups = {'small': [], 'medium': [], 'large': []}\n        self.global_mean = None\n        self.global_std = None\n        \n    def fit(self, X, y):\n        \"\"\"\n        Fit the model using reference structures\n        \n        Parameters:\n        - X: Features (one-hot encoded RNA sequences)\n        - y: 3D coordinates of reference structures\n        \"\"\"\n        n_samples = len(X)\n        self.reference_structures = []\n        self.reference_sequences = []\n        \n        # Extract sequences from one-hot encoding\n        nucleotides = ['A', 'C', 'G', 'U', 'N']\n        \n        all_coords = []\n        \n        for i in range(n_samples):\n            # Get the sequence length by finding the first all-zero row\n            seq_length = X[i].shape[0]\n            for j in range(X[i].shape[0]):\n                if np.all(X[i][j] == 0):\n                    seq_length = j\n                    break\n            \n            # Extract the sequence\n            sequence = ''\n            for j in range(seq_length):\n                idx = np.argmax(X[i][j])\n                if idx < len(nucleotides):\n                    sequence += nucleotides[idx]\n                else:\n                    sequence += 'N'\n            \n            # Get coordinates for this sequence\n            coords = y[i][:seq_length]\n            \n            # Store valid reference structures\n            if not np.isnan(coords).any() and not np.all(coords == 0):\n                self.reference_structures.append(coords)\n                self.reference_sequences.append(sequence)\n                \n                # Categorize by size\n                if seq_length < 60:\n                    self.size_groups['small'].append(len(self.reference_structures) - 1)\n                elif seq_length < 200:\n                    self.size_groups['medium'].append(len(self.reference_structures) - 1)\n                else:\n                    self.size_groups['large'].append(len(self.reference_structures) - 1)\n                \n                all_coords.append(coords)\n        \n        # Calculate global statistics\n        if all_coords:\n            all_coords_array = np.vstack(all_coords)\n            self.global_mean = np.mean(all_coords_array, axis=0)\n            self.global_std = np.std(all_coords_array, axis=0)\n        else:\n            self.global_mean = np.zeros(3)\n            self.global_std = np.ones(3) * 10  # Default standard deviation\n            \n        print(f\"Model fit with {len(self.reference_structures)} reference structures\")\n        print(f\"Size groups: Small={len(self.size_groups['small'])}, \"\n              f\"Medium={len(self.size_groups['medium'])}, \"\n              f\"Large={len(self.size_groups['large'])}\")\n    \n    def predict(self, X):\n        \"\"\"\n        Predict 3D structures for RNA sequences\n        \n        Parameters:\n        - X: Features (one-hot encoded RNA sequences)\n        \n        Returns:\n        - Predicted 3D coordinates\n        \"\"\"\n        n_samples = len(X)\n        predictions = []\n        \n        nucleotides = ['A', 'C', 'G', 'U', 'N']\n        \n        for i in range(n_samples):\n            # Get the sequence length by finding the first all-zero row\n            seq_length = X[i].shape[0]\n            for j in range(X[i].shape[0]):\n                if np.all(X[i][j] == 0):\n                    seq_length = j\n                    break\n            \n            # Extract the sequence\n            sequence = ''\n            for j in range(seq_length):\n                if j < X[i].shape[0]:\n                    idx = np.argmax(X[i][j])\n                    if idx < len(nucleotides):\n                        sequence += nucleotides[idx]\n                    else:\n                        sequence += 'N'\n                else:\n                    break\n            \n            # Find MSA file if available\n            msa_file = None\n            if self.msa_dir is not None:\n                # Extract sequence ID from one-hot encoded representation\n                # This requires knowledge of how the sequence IDs are stored\n                target_id = f\"seq_{i+1}\"  # Default fallback\n                \n                # Look for MSA file\n                potential_msa_file = os.path.join(self.msa_dir, f\"{target_id}.MSA.fasta\")\n                if os.path.exists(potential_msa_file):\n                    msa_file = potential_msa_file\n            \n            # Adjust noise level based on sequence length\n            if seq_length < 60:\n                group = \"small\"\n                noise_level = self.base_noise_level * 1.5\n            elif seq_length < 200:\n                group = \"medium\"\n                noise_level = self.base_noise_level * 1.0\n            else:\n                group = \"large\"\n                noise_level = self.base_noise_level * 0.6\n            \n            # Create a prediction for this sequence\n            base_struct = None\n            \n            # If we have reference structures in this size group, use them\n            if group in self.size_groups and self.size_groups[group]:\n                # Find the most similar sequence in terms of length\n                best_idx = -1\n                best_length_diff = float('inf')\n                \n                for idx in self.size_groups[group]:\n                    ref_seq = self.reference_sequences[idx]\n                    length_diff = abs(len(ref_seq) - seq_length)\n                    \n                    if length_diff < best_length_diff:\n                        best_length_diff = length_diff\n                        best_idx = idx\n                \n                if best_idx >= 0 and best_length_diff <= seq_length * 0.5:  # Only use if reasonably close\n                    base_struct = self.reference_structures[best_idx].copy()\n                    \n                    # If reference is shorter, extend it\n                    if len(base_struct) < seq_length:\n                        extension = np.zeros((seq_length - len(base_struct), 3))\n                        # Extrapolate the last few positions to extend\n                        if len(base_struct) > 2:\n                            direction = base_struct[-1] - base_struct[-2]\n                            for j in range(seq_length - len(base_struct)):\n                                extension[j] = base_struct[-1] + direction * (j + 1)\n                        base_struct = np.vstack([base_struct, extension])\n                    \n                    # If reference is longer, truncate it\n                    if len(base_struct) > seq_length:\n                        base_struct = base_struct[:seq_length]\n                    \n                    # Apply variation\n                    if self.geometric_sampling:\n                        pred = enhanced_structural_variation(\n                            base_struct, \n                            noise_level=noise_level,\n                            preserve_distance=True,\n                            use_global_movement=(group == \"small\"),\n                            correlation=self.correlation,\n                            sequence=sequence\n                        )\n                    else:\n                        noise = np.random.normal(0, noise_level, base_struct.shape)\n                        pred = base_struct + noise\n            \n            # Fall back to a generic structure if no suitable reference found\n            if base_struct is None:\n                # Create a simple extended chain as starting structure\n                pred = np.zeros((seq_length, 3))\n                for j in range(seq_length):\n                    pred[j] = np.array([j * ADJACENT_NUCLEOTIDE_DISTANCE, 0, 0])\n                \n                # Apply variation to make it more realistic\n                if self.geometric_sampling:\n                    pred = enhanced_structural_variation(\n                        pred, \n                        noise_level=noise_level * 2,  # Higher noise for generic structures\n                        preserve_distance=True,\n                        use_global_movement=True,\n                        correlation=self.correlation,\n                        sequence=sequence\n                    )\n                else:\n                    noise = np.random.normal(0, noise_level * 2, pred.shape)\n                    pred = pred + noise\n            \n            # Apply base pairing constraints\n            pred = apply_base_pairing_constraints(pred, sequence)\n            \n            # Apply MSA constraints if available\n            if msa_file:\n                pred = apply_msa_constraints(pred, sequence, msa_file)\n            \n            # Pad with zeros to match maximum length\n            if len(pred) < X[i].shape[0]:\n                padded_pred = np.zeros((X[i].shape[0], 3))\n                padded_pred[:len(pred)] = pred\n                pred = padded_pred\n            \n            predictions.append(pred)\n        \n        return np.array(predictions)\n    \n    def predict_ensemble(self, X, test_seq_df=None):\n        \"\"\"\n        Generate an ensemble of predictions for each sequence\n        \n        Parameters:\n        - X: Features (one-hot encoded RNA sequences)\n        - test_seq_df: DataFrame with sequence information (for MSA lookup)\n        \n        Returns:\n        - Dictionary mapping sequence IDs to lists of structures\n        \"\"\"\n        n_samples = len(X)\n        ensemble_predictions = {}\n        \n        nucleotides = ['A', 'C', 'G', 'U', 'N']\n        \n        for i in range(n_samples):\n            # Extract target_id if available\n            target_id = f\"seq_{i+1}\"  # Default\n            if test_seq_df is not None and i < len(test_seq_df):\n                target_id = test_seq_df.iloc[i]['target_id']\n            \n            # Get the sequence length\n            seq_length = X[i].shape[0]\n            for j in range(X[i].shape[0]):\n                if np.all(X[i][j] == 0):\n                    seq_length = j\n                    break\n            \n            # Extract the sequence\n            sequence = ''\n            for j in range(seq_length):\n                if j < X[i].shape[0]:\n                    idx = np.argmax(X[i][j])\n                    if idx < len(nucleotides):\n                        sequence += nucleotides[idx]\n                    else:\n                        sequence += 'N'\n                else:\n                    break\n            \n            # Find MSA file if available\n            msa_file = None\n            if self.msa_dir is not None:\n                potential_msa_files = [\n                    os.path.join(self.msa_dir, f\"{target_id}.MSA.fasta\"),\n                    os.path.join(self.msa_dir, f\"{target_id.split('_')[0]}.MSA.fasta\")\n                ]\n                \n                for potential_file in potential_msa_files:\n                    if os.path.exists(potential_file):\n                        msa_file = potential_file\n                        break\n            \n            # Get base prediction\n            base_prediction = self.predict([X[i]])[0][:seq_length]\n            \n            # Generate diverse ensemble\n            structures = generate_diverse_ensemble(\n                base_prediction, \n                sequence, \n                msa_file=msa_file\n            )\n            \n            # Store all structures for this sequence\n            ensemble_predictions[target_id] = structures\n        \n        return ensemble_predictions ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:07:19.659148Z","iopub.execute_input":"2025-04-19T16:07:19.659519Z","iopub.status.idle":"2025-04-19T16:07:19.675620Z","shell.execute_reply.started":"2025-04-19T16:07:19.659490Z","shell.execute_reply":"2025-04-19T16:07:19.675007Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport time\nimport gc\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport warnings\nimport traceback\nfrom tqdm import tqdm\nfrom Bio import PDB\nfrom Bio.PDB import PDBParser\nfrom Bio.PDB.MMCIF2Dict import MMCIF2Dict\nimport urllib.request\nimport gzip\nimport shutil\nimport requests\nfrom io import StringIO\nimport json\nimport pickle\nimport sys\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:07:29.798516Z","iopub.execute_input":"2025-04-19T16:07:29.798796Z","iopub.status.idle":"2025-04-19T16:07:29.932388Z","shell.execute_reply.started":"2025-04-19T16:07:29.798772Z","shell.execute_reply":"2025-04-19T16:07:29.931528Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Import our improved model\nfrom model import (\n    ImprovedRNAModel, \n    enhanced_structural_variation,\n    apply_base_pairing_constraints,\n    apply_msa_constraints,\n    normalize_structure,\n    generate_diverse_ensemble,\n    ADJACENT_NUCLEOTIDE_DISTANCE,\n    BASE_PAIR_DISTANCE\n)\n\n# Suppress warnings\nwarnings.filterwarnings('ignore')\n\n# Set random seed for reproducibility\nnp.random.seed(42)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:07:37.019611Z","iopub.execute_input":"2025-04-19T16:07:37.020300Z","iopub.status.idle":"2025-04-19T16:07:37.030011Z","shell.execute_reply.started":"2025-04-19T16:07:37.020276Z","shell.execute_reply":"2025-04-19T16:07:37.029441Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Directories and files\nDATA_DIR = os.getenv('DATA_DIR', '/kaggle/input/stanford-rna-3d-folding/')\nOUTPUT_DIR = os.getenv('OUTPUT_DIR', '/kaggle/working/')\nMSA_DIR = os.path.join(DATA_DIR, \"MSA\")\nEXTERNAL_DATA_DIR = os.path.join(OUTPUT_DIR, \"external_data\")\nCACHE_DIR = os.path.join(OUTPUT_DIR, \"cache\")\n\n# Add path to pre-downloaded RNA data\nRNA_DATA_DIR = \"rna_data\"\nRNA_STRUCTURES_DIR = os.path.join(RNA_DATA_DIR, \"structures\")\nRNA_CHAINS_CACHE = os.path.join(RNA_DATA_DIR, \"rna_chains.pkl\")\n\n# Create necessary directories\nos.makedirs(OUTPUT_DIR, exist_ok=True)\nos.makedirs(EXTERNAL_DATA_DIR, exist_ok=True)\nos.makedirs(CACHE_DIR, exist_ok=True)\n\n# Local cache for RNA structures\nRNA_STRUCTURES_CACHE = os.path.join(CACHE_DIR, \"rna_structures.pkl\")\n\n# Hardcoded RNA PDB IDs as fallback\nFALLBACK_PDB_IDS = [\n    \"6th6\", \"6adr\", \"6qir\", \"6vwl\", \"6zu0\", \"7dk6\", \"7eag\",\n    \"7egp\", \"4u3l\", \"4u3k\", \"1bgz\", \"1eh4\", \"1eht\", \"1equ\",\n    \"1exd\", \"1f7y\", \"1jid\", \"1kd1\", \"1kxk\", \"1l8v\", \"1o0c\",\n    \"1r3o\", \"1s72\", \"1x8w\", \"1y27\", \"2oe5\", \"2zni\", \"3egz\",\n    \"3g78\", \"3gx5\", \"3hxm\", \"3j0o\", \"3j0p\", \"3j0q\", \"3j0r\"\n]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:07:49.237912Z","iopub.execute_input":"2025-04-19T16:07:49.238462Z","iopub.status.idle":"2025-04-19T16:07:49.244574Z","shell.execute_reply.started":"2025-04-19T16:07:49.238442Z","shell.execute_reply":"2025-04-19T16:07:49.243950Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Hard-coded RNA structures as fallback (pre-parsed chains)\nHARDCODED_RNA_CHAINS = [\n    # A simplified representation of RNA chains with sequence and coordinates\n    # These will be used if no internet connection is available\n    # Format: {'chain_id': 'A', 'sequence': 'ACGU...', 'coordinates': numpy array of shape (seq_len, 3)}\n]\n\n# Local RNA structure files (add your own PDB files here)\nLOCAL_RNA_STRUCTURES = [\n    # Add paths to local RNA structure files if available\n]\n\n# Add paths to external Ribonanza data files\nRIBONANZA_SEQ_FILE = \"/kaggle/input/parquet-files-for-stanford/ext_ribonanza_labels.parquet\"\nRIBONANZA_LABELS_FILE = \"/kaggle/input/parquet-files-for-stanford/ext_ribonanza_sequences.parquet\"\n\n# Set offline mode as default\nos.environ['OFFLINE_MODE'] = 'True'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:08:18.916687Z","iopub.execute_input":"2025-04-19T16:08:18.917246Z","iopub.status.idle":"2025-04-19T16:08:18.921181Z","shell.execute_reply.started":"2025-04-19T16:08:18.917221Z","shell.execute_reply":"2025-04-19T16:08:18.920368Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===== DEPENDENCY MANAGEMENT =====\ndef check_and_install_dependencies():\n    \"\"\"\n    Check if required packages are installed and install them locally if needed\n    \"\"\"\n    required_packages = {\n        'numpy': 'numpy',\n        'pandas': 'pandas',\n        'matplotlib': 'matplotlib',\n        'tqdm': 'tqdm',\n        'biopython': 'Bio',\n        'requests': 'requests'\n    }\n    \n    missing_packages = []\n    for package_name, import_name in required_packages.items():\n        try:\n            __import__(import_name)\n            print(f\"✓ {package_name} is installed\")\n        except ImportError:\n            missing_packages.append(package_name)\n            print(f\"✗ {package_name} is missing\")\n    \n    if missing_packages:\n        print(\"\\nMissing packages detected. Please install them before running:\")\n        for package in missing_packages:\n            print(f\"pip install {package}\")\n        print(\"\\nFor offline submission, download packages with:\")\n        print(\"pip download -d ./packages numpy pandas matplotlib tqdm biopython requests\")\n        print(\"Then install offline with:\")\n        print(\"pip install --no-index --find-links=./packages -r requirements.txt\")\n        \n        # Create requirements.txt file\n        with open(\"requirements.txt\", \"w\") as f:\n            for package in required_packages.keys():\n                f.write(f\"{package}\\n\")\n        \n        print(\"\\nrequirements.txt has been created.\")\n        return False\n    \n    return True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:08:21.478537Z","iopub.execute_input":"2025-04-19T16:08:21.479090Z","iopub.status.idle":"2025-04-19T16:08:21.484497Z","shell.execute_reply.started":"2025-04-19T16:08:21.479062Z","shell.execute_reply":"2025-04-19T16:08:21.483759Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===== EXTERNAL DATA LOADING =====\ndef download_pdb_structure(pdb_id, output_dir=EXTERNAL_DATA_DIR):\n    \"\"\"\n    Download PDB structure from RCSB PDB database - DISABLED FOR OFFLINE USE\n    \"\"\"\n    print(f\"Offline mode: Cannot download {pdb_id} from internet\")\n    return None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:08:28.473611Z","iopub.execute_input":"2025-04-19T16:08:28.473892Z","iopub.status.idle":"2025-04-19T16:08:28.477926Z","shell.execute_reply.started":"2025-04-19T16:08:28.473870Z","shell.execute_reply":"2025-04-19T16:08:28.477062Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def download_rna_structures_from_rna3dhub(output_dir=EXTERNAL_DATA_DIR, max_structures=100):\n    \"\"\"\n    Download RNA structures from RNA 3D Hub - DISABLED FOR OFFLINE USE\n    \"\"\"\n    print(\"Offline mode: Cannot download RNA structures from internet\")\n    return []","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:08:39.626531Z","iopub.execute_input":"2025-04-19T16:08:39.626791Z","iopub.status.idle":"2025-04-19T16:08:39.630728Z","shell.execute_reply.started":"2025-04-19T16:08:39.626771Z","shell.execute_reply":"2025-04-19T16:08:39.630018Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def parse_pdb_structure(structure_file):\n    \"\"\"\n    Parse a PDB/mmCIF file and extract RNA chain information\n    \"\"\"\n    try:\n        parser = PDBParser(QUIET=True)\n        structure = None\n        \n        # Determine file type and parse accordingly\n        if structure_file.endswith('.pdb'):\n            structure = parser.get_structure('RNA', structure_file)\n        elif structure_file.endswith('.cif'):\n            # Use MMCIF parser for cif files\n            from Bio.PDB.MMCIFParser import MMCIFParser\n            mmcif_parser = MMCIFParser(QUIET=True)\n            structure = mmcif_parser.get_structure('RNA', structure_file)\n        else:\n            print(f\"Unsupported file format for {structure_file}\")\n            return []\n        \n        if structure is None:\n            return []\n        \n        # Extract RNA chains\n        rna_chains = []\n        \n        for model in structure:\n            for chain in model:\n                # Check if this is an RNA chain\n                is_rna = False\n                sequence = \"\"\n                coordinates = []\n                \n                for residue in chain:\n                    # Check if it's a nucleotide (contains C1' atom)\n                    if 'C1\\'' in residue:\n                        is_rna = True\n                        # Extract the nucleotide identity\n                        if residue.get_resname() in ['A', 'C', 'G', 'U']:\n                            nucleotide = residue.get_resname()\n                        elif residue.get_resname() in ['DA', 'DC', 'DG', 'DT']:\n                            # Convert DNA to RNA\n                            dna_to_rna = {'DA': 'A', 'DC': 'C', 'DG': 'G', 'DT': 'U'}\n                            nucleotide = dna_to_rna.get(residue.get_resname(), 'N')\n                        else:\n                            # Handle other nucleotide naming conventions\n                            if 'A' in residue.get_resname():\n                                nucleotide = 'A'\n                            elif 'C' in residue.get_resname():\n                                nucleotide = 'C'\n                            elif 'G' in residue.get_resname():\n                                nucleotide = 'G'\n                            elif 'U' in residue.get_resname() or 'T' in residue.get_resname():\n                                nucleotide = 'U'\n                            else:\n                                nucleotide = 'N'\n                        \n                        sequence += nucleotide\n                        \n                        # Get C1' atom coordinates\n                        try:\n                            c1_atom = residue['C1\\'']\n                            coordinates.append([c1_atom.get_coord()[0], \n                                              c1_atom.get_coord()[1], \n                                              c1_atom.get_coord()[2]])\n                        except KeyError:\n                            # If C1' is not available, try to use another atom\n                            for atom in residue:\n                                if atom.get_name() in ['P', 'C4\\'']:\n                                    coordinates.append([atom.get_coord()[0], \n                                                      atom.get_coord()[1], \n                                                      atom.get_coord()[2]])\n                                    break\n                            else:\n                                # If no suitable atom found, use average of all atoms\n                                all_coords = np.array([atom.get_coord() for atom in residue])\n                                avg_coord = np.mean(all_coords, axis=0)\n                                coordinates.append([avg_coord[0], avg_coord[1], avg_coord[2]])\n                \n                if is_rna and len(sequence) >= 10:  # Only consider chains with at least 10 nucleotides\n                    rna_chains.append({\n                        'chain_id': chain.get_id(),\n                        'sequence': sequence,\n                        'coordinates': np.array(coordinates)\n                    })\n        \n        return rna_chains\n    \n    except Exception as e:\n        print(f\"Error parsing structure file {structure_file}: {str(e)}\")\n        return []\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:08:41.241319Z","iopub.execute_input":"2025-04-19T16:08:41.241776Z","iopub.status.idle":"2025-04-19T16:08:41.251563Z","shell.execute_reply.started":"2025-04-19T16:08:41.241753Z","shell.execute_reply":"2025-04-19T16:08:41.250789Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_external_rna_structures(max_structures=100):\n    \"\"\"\n    Load external RNA structures from local files or generate synthetic data\n    \"\"\"\n    print(\"Loading external RNA structures...\")\n    \n    # First check if we have pre-downloaded data from prepare_rna_data.py\n    if os.path.exists(RNA_CHAINS_CACHE):\n        try:\n            print(f\"Loading pre-downloaded RNA chains from: {RNA_CHAINS_CACHE}\")\n            with open(RNA_CHAINS_CACHE, 'rb') as f:\n                rna_chains = pickle.load(f)\n            print(f\"Loaded {len(rna_chains)} pre-downloaded RNA chains\")\n            return rna_chains\n        except Exception as e:\n            print(f\"Error loading pre-downloaded RNA chains: {str(e)}\")\n    # Check if we have cached structures from previous runs\n    if os.path.exists(RNA_STRUCTURES_CACHE):\n        try:\n            print(f\"Loading RNA structures from cache: {RNA_STRUCTURES_CACHE}\")\n            with open(RNA_STRUCTURES_CACHE, 'rb') as f:\n                rna_chains = pickle.load(f)\n            print(f\"Loaded {len(rna_chains)} RNA chains from cache\")\n            return rna_chains\n        except Exception as e:\n            print(f\"Error loading cached structures: {str(e)}\")\n    \n    # Check for pre-downloaded structure files\n    pre_downloaded_files = []\n    if os.path.exists(RNA_STRUCTURES_DIR):\n        pre_downloaded_files = [\n            os.path.join(RNA_STRUCTURES_DIR, f) \n            for f in os.listdir(RNA_STRUCTURES_DIR) \n            if f.endswith(('.pdb', '.cif'))\n        ]\n        if pre_downloaded_files:\n            print(f\"Found {len(pre_downloaded_files)} pre-downloaded structure files\")\n    \n    # Try to load pre-downloaded structure files\n    if pre_downloaded_files:\n        print(f\"Loading {len(pre_downloaded_files)} pre-downloaded structure files\")\n        rna_chains = []\n        for structure_file in tqdm(pre_downloaded_files, desc=\"Parsing pre-downloaded structures\"):\n            chains = parse_pdb_structure(structure_file)\n            rna_chains.extend(chains)\n        \n        print(f\"Loaded {len(rna_chains)} RNA chains from pre-downloaded structures\")\n        \n        # Cache the results\n        if len(rna_chains) > 0:\n            try:\n                with open(RNA_STRUCTURES_CACHE, 'wb') as f:\n                    pickle.dump(rna_chains, f)\n                print(f\"Cached {len(rna_chains)} RNA chains for future use\")\n            except Exception as e:\n                print(f\"Error caching RNA chains: {str(e)}\")\n            \n        return rna_chains\n    elif HARDCODED_RNA_CHAINS:\n        print(\"Using hardcoded RNA structures\")\n        return HARDCODED_RNA_CHAINS\n    else:\n        # Generate synthetic RNA structures as fallback\n        print(\"Generating synthetic RNA structures as fallback\")\n        return generate_synthetic_rna_structures(20)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:08:44.690214Z","iopub.execute_input":"2025-04-19T16:08:44.690467Z","iopub.status.idle":"2025-04-19T16:08:44.700258Z","shell.execute_reply.started":"2025-04-19T16:08:44.690448Z","shell.execute_reply":"2025-04-19T16:08:44.699539Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def generate_synthetic_rna_structures(num_structures=20):\n    \"\"\"\n    Generate synthetic RNA structures as fallback\n    \"\"\"\n    rna_chains = []\n    \n    for i in range(num_structures):\n        # Generate random RNA sequence of length 30-100\n        length = np.random.randint(30, 101)\n        nucleotides = ['A', 'C', 'G', 'U']\n        sequence = ''.join(np.random.choice(nucleotides) for _ in range(length))\n        \n        # Generate random 3D coordinates in a realistic range\n        # Start with a linear chain\n        coords = np.zeros((length, 3))\n        for j in range(length):\n            # Add some randomness to the linear structure\n            if j == 0:\n                coords[j] = np.random.normal(0, 1, 3)\n            else:\n                # Average distance between adjacent nucleotides ~6Å\n                direction = np.random.normal(0, 1, 3)\n                direction = direction / np.linalg.norm(direction) * 6.0\n                coords[j] = coords[j-1] + direction\n        \n        rna_chains.append({\n            'chain_id': f'synthetic_{i}',\n            'sequence': sequence,\n            'coordinates': coords\n        })\n    \n    return rna_chains\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:08:47.057301Z","iopub.execute_input":"2025-04-19T16:08:47.057903Z","iopub.status.idle":"2025-04-19T16:08:47.063989Z","shell.execute_reply.started":"2025-04-19T16:08:47.057875Z","shell.execute_reply":"2025-04-19T16:08:47.063326Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def convert_external_data_to_features(rna_chains, max_length=720):\n    \"\"\"\n    Convert external RNA structures to feature format compatible with our model\n    \"\"\"\n    if not rna_chains:\n        print(\"No external RNA chains available, skipping conversion\")\n        return np.array([]), np.array([])\n        \n    X_external = []\n    y_external = []\n    \n    for chain in rna_chains:\n        sequence = chain['sequence']\n        coordinates = chain['coordinates']\n        \n        # Skip if sequence is too long or coordinates don't match\n        if len(sequence) > max_length or len(sequence) != len(coordinates):\n            continue\n        \n        # Convert sequence to one-hot encoded features\n        features = []\n        for nucleotide in sequence:\n            if nucleotide == 'A':\n                features.append([1, 0, 0, 0, 0])\n            elif nucleotide == 'C':\n                features.append([0, 1, 0, 0, 0])\n            elif nucleotide == 'G':\n                features.append([0, 0, 1, 0, 0])\n            elif nucleotide == 'U':\n                features.append([0, 0, 0, 1, 0])\n            else:\n                features.append([0, 0, 0, 0, 1])\n        \n        # Pad or truncate features\n        if len(features) < max_length:\n            padding = [[0, 0, 0, 0, 0]] * (max_length - len(features))\n            features.extend(padding)\n        else:\n            features = features[:max_length]\n        \n        # Pad or truncate coordinates\n        coords = np.zeros((max_length, 3))\n        coords[:len(coordinates)] = coordinates\n        \n        X_external.append(features)\n        y_external.append(coords)\n    \n    return np.array(X_external), np.array(y_external)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:08:49.743324Z","iopub.execute_input":"2025-04-19T16:08:49.743604Z","iopub.status.idle":"2025-04-19T16:08:49.750683Z","shell.execute_reply.started":"2025-04-19T16:08:49.743582Z","shell.execute_reply":"2025-04-19T16:08:49.749869Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===== COMPETITION DATA LOADING =====\ndef load_competition_data():\n    \"\"\"\n    Load main data files for the competition.\n    \"\"\"\n    print(\"Loading competition data...\")\n    main_files = {\n        \"train_sequences\": os.path.join(DATA_DIR, \"train_sequences.csv\"),\n        \"train_labels\": os.path.join(DATA_DIR, \"train_labels.csv\"),\n        \"validation_sequences\": os.path.join(DATA_DIR, \"validation_sequences.csv\"),\n        \"validation_labels\": os.path.join(DATA_DIR, \"validation_labels.csv\"),\n        \"test_sequences\": os.path.join(DATA_DIR, \"test_sequences.csv\"),\n        \"sample_submission\": os.path.join(DATA_DIR, \"sample_submission.csv\")\n    }\n    \n    data = {}\n    for name, file_path in main_files.items():\n        if os.path.exists(file_path):\n            data[name] = pd.read_csv(file_path)\n            print(f\"Loaded {name}: {data[name].shape}\")\n        else:\n            print(f\"Warning: File {file_path} not found\")\n            data[name] = None\n    \n    return data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:08:52.230328Z","iopub.execute_input":"2025-04-19T16:08:52.230905Z","iopub.status.idle":"2025-04-19T16:08:52.235829Z","shell.execute_reply.started":"2025-04-19T16:08:52.230885Z","shell.execute_reply":"2025-04-19T16:08:52.235159Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def prepare_features(seq_df, max_length=720):\n    \"\"\"\n    Prepare one-hot encoded features from RNA sequences.\n    \"\"\"\n    print(f\"Preparing features for {len(seq_df)} sequences...\")\n    X = []\n    for _, row in seq_df.iterrows():\n        seq = row['sequence']\n        features = []\n        for nucleotide in seq:\n            if nucleotide == 'A':\n                features.append([1, 0, 0, 0, 0])\n            elif nucleotide == 'C':\n                features.append([0, 1, 0, 0, 0])\n            elif nucleotide == 'G':\n                features.append([0, 0, 1, 0, 0])\n            elif nucleotide == 'U':\n                features.append([0, 0, 0, 1, 0])\n            else:\n                features.append([0, 0, 0, 0, 1])\n        \n        # Pad or truncate to max_length\n        if len(features) < max_length:\n            padding = [[0, 0, 0, 0, 0]] * (max_length - len(features))\n            features.extend(padding)\n        else:\n            features = features[:max_length]\n        \n        X.append(features)\n    \n    return np.array(X)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:08:54.245633Z","iopub.execute_input":"2025-04-19T16:08:54.246320Z","iopub.status.idle":"2025-04-19T16:08:54.252545Z","shell.execute_reply.started":"2025-04-19T16:08:54.246295Z","shell.execute_reply":"2025-04-19T16:08:54.251717Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def prepare_labels(labels_df, seq_df, max_length=720):\n    \"\"\"\n    Prepare 3D coordinate labels from label dataframe.\n    \"\"\"\n    print(f\"Preparing labels for {len(seq_df)} sequences...\")\n    y = []\n    \n    # Group labels by sequence\n    for i, row in seq_df.iterrows():\n        target_id = row['target_id']\n        seq_length = len(row['sequence'])\n        \n        # Extract all residues for this sequence\n        seq_labels = labels_df[labels_df['ID'].str.startswith(f\"{target_id}_\")]\n        \n        # Create coordinates array\n        coords = np.zeros((max_length, 3))\n        \n        if not seq_labels.empty:\n            for _, label_row in seq_labels.iterrows():\n                resid = int(label_row['resid'])\n                if resid <= max_length:\n                    coords[resid-1, 0] = label_row['x_1']\n                    coords[resid-1, 1] = label_row['y_1']\n                    coords[resid-1, 2] = label_row['z_1']\n        \n        y.append(coords)\n    \n    return np.array(y)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:08:57.063000Z","iopub.execute_input":"2025-04-19T16:08:57.063551Z","iopub.status.idle":"2025-04-19T16:08:57.068996Z","shell.execute_reply.started":"2025-04-19T16:08:57.063529Z","shell.execute_reply":"2025-04-19T16:08:57.068240Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===== ADDITIONAL EXTERNAL DATA LOADING =====\ndef load_ribonanza_data(max_length=720):\n    \"\"\"\n    Load external Ribonanza data from parquet files\n    \"\"\"\n    print(\"Loading Ribonanza external data...\")\n    \n    # Check if files exist\n    if not os.path.exists(RIBONANZA_SEQ_FILE) or not os.path.exists(RIBONANZA_LABELS_FILE):\n        print(\"Ribonanza parquet files not found, skipping this data source\")\n        return np.array([]), np.array([])\n    \n    try:\n        # Load the data\n        sequences_df = pd.read_parquet(RIBONANZA_SEQ_FILE)\n        labels_df = pd.read_parquet(RIBONANZA_LABELS_FILE)\n        \n        print(f\"Loaded Ribonanza data: {len(sequences_df)} sequences\")\n        \n        # Prepare features and labels\n        X_ribo = []\n        y_ribo = []\n        \n        for _, row in tqdm(sequences_df.iterrows(), total=len(sequences_df), desc=\"Processing Ribonanza sequences\"):\n            seq_id = row.get('sequence_id', row.get('ID', None))\n            if seq_id is None:\n                continue\n                \n            sequence = row.get('sequence', '')\n            if not sequence:\n                continue\n            \n            # Get coordinates for this sequence\n            coords = labels_df[labels_df['sequence_id'] == seq_id]\n            if coords.empty:\n                continue\n            \n            # Convert sequence to one-hot encoded features\n            features = []\n            for nucleotide in sequence:\n                if nucleotide == 'A':\n                    features.append([1, 0, 0, 0, 0])\n                elif nucleotide == 'C':\n                    features.append([0, 1, 0, 0, 0])\n                elif nucleotide == 'G':\n                    features.append([0, 0, 1, 0, 0])\n                elif nucleotide == 'U':\n                    features.append([0, 0, 0, 1, 0])\n                else:\n                    features.append([0, 0, 0, 0, 1])\n            \n            # Check if sequence is too long\n            if len(features) > max_length:\n                print(f\"Sequence {seq_id} is too long ({len(features)} > {max_length}), truncating\")\n                features = features[:max_length]\n            \n            # Pad if needed\n            if len(features) < max_length:\n                padding = [[0, 0, 0, 0, 0]] * (max_length - len(features))\n                features.extend(padding)\n            \n            # Process coordinates\n            coordinates = np.zeros((max_length, 3))\n            \n            # Check which columns contain 3D coordinates\n            coord_columns = []\n            for col in coords.columns:\n                if col.startswith('x_') or col.startswith('y_') or col.startswith('z_'):\n                    coord_columns.append(col)\n            \n            if not coord_columns:\n                # Try to look for other coordinate column patterns\n                x_cols = [col for col in coords.columns if 'x' in col.lower()]\n                y_cols = [col for col in coords.columns if 'y' in col.lower()]\n                z_cols = [col for col in coords.columns if 'z' in col.lower()]\n                \n                if x_cols and y_cols and z_cols:\n                    # Use first set of coordinate columns found\n                    for i, row in coords.iterrows():\n                        pos = min(int(i), max_length-1)\n                        coordinates[pos, 0] = row[x_cols[0]]\n                        coordinates[pos, 1] = row[y_cols[0]]\n                        coordinates[pos, 2] = row[z_cols[0]]\n            else:\n                # Use standard coordinate columns\n                for i, row in coords.iterrows():\n                    position = min(int(i), max_length-1)\n                    x_col = next((col for col in coord_columns if col.startswith('x_')), None)\n                    y_col = next((col for col in coord_columns if col.startswith('y_')), None)\n                    z_col = next((col for col in coord_columns if col.startswith('z_')), None)\n                    \n                    if x_col and y_col and z_col:\n                        coordinates[position, 0] = row[x_col]\n                        coordinates[position, 1] = row[y_col]\n                        coordinates[position, 2] = row[z_col]\n            \n            # Skip if all coordinates are zero (no actual data)\n            if np.all(coordinates == 0):\n                continue\n                \n            X_ribo.append(features)\n            y_ribo.append(coordinates)\n        \n        print(f\"Prepared {len(X_ribo)} Ribonanza sequences with 3D coordinates\")\n        return np.array(X_ribo), np.array(y_ribo)\n        \n    except Exception as e:\n        print(f\"Error loading Ribonanza data: {str(e)}\")\n        traceback.print_exc()\n        return np.array([]), np.array([])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:08:58.883997Z","iopub.execute_input":"2025-04-19T16:08:58.884263Z","iopub.status.idle":"2025-04-19T16:08:58.897146Z","shell.execute_reply.started":"2025-04-19T16:08:58.884242Z","shell.execute_reply":"2025-04-19T16:08:58.896375Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\ndef load_all_data(max_length=720, use_external=True, max_external_structures=100):\n    \"\"\"\n    Load and process all data including external structures if specified.\n    \"\"\"\n    # Load competition data\n    competition_data = load_competition_data()\n    \n    # Prepare competition training data\n    X_train = prepare_features(competition_data[\"train_sequences\"], max_length)\n    y_train = prepare_labels(competition_data[\"train_labels\"], competition_data[\"train_sequences\"], max_length)\n    \n    # Prepare competition validation data\n    X_valid = prepare_features(competition_data[\"validation_sequences\"], max_length)\n    y_valid = prepare_labels(competition_data[\"validation_labels\"], competition_data[\"validation_sequences\"], max_length)\n    \n    # Load external data if specified\n    if use_external:\n        print(\"\\nLoading external RNA structure data...\")\n        print(\"Running in OFFLINE MODE - only local data will be used\")\n        \n        # Load RNA 3D Hub and PDB data from local files only\n        rna_chains = load_external_rna_structures(max_structures=max_external_structures)\n        X_external, y_external = convert_external_data_to_features(rna_chains, max_length)\n        \n        print(f\"External RNA 3D Hub data: {len(X_external)} structures\")\n        \n        # Load Ribonanza data\n        X_ribo, y_ribo = load_ribonanza_data(max_length)\n        print(f\"External Ribonanza data: {len(X_ribo)} structures\")\n        \n        # Combine all external data sources\n        all_external_data = []\n        \n        if len(X_external) > 0:\n            all_external_data.append((X_external, y_external))\n            \n        if len(X_ribo) > 0:\n            all_external_data.append((X_ribo, y_ribo))\n        \n        # Combine competition and external data if we have external data\n        if all_external_data:\n            X_train_combined = X_train\n            y_train_combined = y_train\n            \n            for X_ext, y_ext in all_external_data:\n                X_train_combined = np.vstack([X_train_combined, X_ext])\n                y_train_combined = np.vstack([y_train_combined, y_ext])\n            \n            print(f\"Combined training data: {X_train_combined.shape}\")\n            \n            return X_train_combined, y_train_combined, X_valid, y_valid, competition_data\n    \n    # Fall back to just competition data if no external data or not requested\n    print(f\"Using only competition data: {X_train.shape}\")\n    return X_train, y_train, X_valid, y_valid, competition_data\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:09:01.813994Z","iopub.execute_input":"2025-04-19T16:09:01.814491Z","iopub.status.idle":"2025-04-19T16:09:01.821359Z","shell.execute_reply.started":"2025-04-19T16:09:01.814468Z","shell.execute_reply":"2025-04-19T16:09:01.820539Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===== MODEL TRAINING & EVALUATION =====\ndef calculate_tm_score_approx(pred, true):\n    \"\"\"\n    Calculate an approximation of TM-score.\n    \n    TM-score measures the similarity of two protein structures.\n    This is a simplified version that doesn't perform optimal alignment.\n    \"\"\"\n    # Filter out padding (zeros)\n    mask = ~np.all(true == 0, axis=1)\n    if not np.any(mask):\n        return 0.0\n    \n    true_filtered = true[mask]\n    pred_filtered = pred[mask]\n    \n    Lref = len(true_filtered)\n    \n    # Calculate d0 scaling factor\n    if Lref >= 30:\n        d0 = 1.24 * np.power(Lref - 15, 1/3) - 1.8\n    else:\n        d0 = 0.5\n    \n    # Calculate distances\n    distances = np.sqrt(np.sum((true_filtered - pred_filtered) ** 2, axis=1))\n    \n    # Calculate TM-score\n    tm_score = np.mean(1.0 / (1.0 + (distances / d0) ** 2))\n    \n    return tm_score\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:09:04.054611Z","iopub.execute_input":"2025-04-19T16:09:04.055144Z","iopub.status.idle":"2025-04-19T16:09:04.060030Z","shell.execute_reply.started":"2025-04-19T16:09:04.055119Z","shell.execute_reply":"2025-04-19T16:09:04.059263Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_model(X_train, y_train, X_valid, y_valid, params=None):\n    \"\"\"\n    Train the improved RNA 3D model.\n    \"\"\"\n    if params is None:\n        params = {\n            'geometric_sampling': True,\n            'base_noise_level': 0.25,  # Reduced noise level for better control\n            'correlation': 0.85\n        }\n    \n    print(f\"Training model with parameters: {params}\")\n    \n    # Initialize and train the model\n    model = ImprovedRNAModel(\n        geometric_sampling=params['geometric_sampling'],\n        base_noise_level=params['base_noise_level'],\n        correlation=params['correlation'],\n        msa_dir=MSA_DIR\n    )\n    \n    # Fit the model\n    model.fit(X_train, y_train)\n    \n    # Validate the model\n    print(\"Evaluating model on validation data...\")\n    y_pred = model.predict(X_valid)\n    \n    # Calculate TM-scores\n    tm_scores = []\n    for i in range(len(X_valid)):\n        tm = calculate_tm_score_approx(y_pred[i], y_valid[i])\n        tm_scores.append(tm)\n    \n    avg_tm_score = np.mean(tm_scores)\n    print(f\"Average TM-score on validation: {avg_tm_score:.4f}\")\n    \n    return model, avg_tm_score, tm_scores","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:09:06.122540Z","iopub.execute_input":"2025-04-19T16:09:06.123247Z","iopub.status.idle":"2025-04-19T16:09:06.128477Z","shell.execute_reply.started":"2025-04-19T16:09:06.123220Z","shell.execute_reply":"2025-04-19T16:09:06.127790Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def run_multi_model_ensemble(X_train, y_train, X_valid, y_valid, test_seq_df, sample_submission_df, num_models=5):\n    \"\"\"\n    Train multiple models with different parameters and create an ensemble.\n    \"\"\"\n    print(f\"Running ensemble with {num_models} models...\")\n    \n    # Optimal parameters based on RNA 3D structure characteristics\n    param_sets = [\n        {'geometric_sampling': True, 'base_noise_level': 0.2, 'correlation': 0.9},  # More stable, higher correlation\n        {'geometric_sampling': True, 'base_noise_level': 0.25, 'correlation': 0.85},\n        {'geometric_sampling': True, 'base_noise_level': 0.3, 'correlation': 0.8},\n        {'geometric_sampling': False, 'base_noise_level': 0.2, 'correlation': 0.9},  # Non-geometric variation\n        {'geometric_sampling': True, 'base_noise_level': 0.15, 'correlation': 0.95},  # Very stable, very high correlation\n    ]\n    \n    # Ensure we have enough parameter sets\n    while len(param_sets) < num_models:\n        param_sets.append({\n            'geometric_sampling': np.random.choice([True, False], p=[0.8, 0.2]),  # Favor geometric sampling\n            'base_noise_level': np.random.uniform(0.15, 0.3),  # Lower noise range\n            'correlation': np.random.uniform(0.8, 0.95)  # Higher correlation range\n        })\n    \n    # Train models\n    models = []\n    scores = []\n    \n    for i, params in enumerate(param_sets[:num_models]):\n        print(f\"\\nTraining model {i+1}/{num_models}\")\n        model, score, _ = train_model(X_train, y_train, X_valid, y_valid, params)\n        models.append(model)\n        scores.append(score)\n    \n    # Prepare test features\n    X_test = prepare_features(test_seq_df)\n    \n    # Generate predictions for each model\n    print(\"\\nGenerating predictions from all models...\")\n    model_predictions = []\n    for i, model in enumerate(models):\n        print(f\"Model {i+1}/{len(models)} (score: {scores[i]:.4f})\")\n        pred = model.predict(X_test)\n        model_predictions.append(pred)\n    \n    # Create submission from ensemble\n    print(\"\\nCreating submission from ensemble...\")\n    return create_ensemble_submission(model_predictions, test_seq_df, sample_submission_df, scores)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:09:08.122895Z","iopub.execute_input":"2025-04-19T16:09:08.123363Z","iopub.status.idle":"2025-04-19T16:09:08.130195Z","shell.execute_reply.started":"2025-04-19T16:09:08.123342Z","shell.execute_reply":"2025-04-19T16:09:08.129521Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_ensemble_submission(model_predictions, test_seq_df, sample_submission_df, model_scores=None):\n    \"\"\"\n    Create a submission file from ensemble predictions with improved weighting.\n    \"\"\"\n    submission_df = sample_submission_df.copy()\n    \n    # Use model scores for weighted averaging if available\n    if model_scores is not None:\n        # Normalize scores to sum to 1\n        weights = np.array(model_scores) / sum(model_scores)\n        print(f\"Using weighted ensemble with weights: {weights}\")\n    else:\n        weights = np.ones(len(model_predictions)) / len(model_predictions)\n    \n    # Dictionary to store structures for each sequence\n    seq_to_structures = {}\n    \n    # Process each test sequence\n    for i, row in tqdm(test_seq_df.iterrows(), total=len(test_seq_df), desc=\"Processing sequences\"):\n        target_id = row['target_id']\n        seq_length = len(row['sequence'])\n        sequence = row['sequence']\n        \n        # Get base predictions from all models for this sequence\n        sequence_preds = [pred[i][:seq_length] for pred in model_predictions]\n        \n        # Apply weights to generate a weighted average prediction\n        weighted_avg = np.zeros_like(sequence_preds[0])\n        for j, pred in enumerate(sequence_preds):\n            weighted_avg += weights[j] * pred\n        \n        # Check if MSA file exists\n        msa_file = None\n        if os.path.exists(os.path.join(MSA_DIR, f\"{target_id}.MSA.fasta\")):\n            msa_file = os.path.join(MSA_DIR, f\"{target_id}.MSA.fasta\")\n        \n        # Generate ensemble of 5 diverse structures using our improved method\n        structures = generate_diverse_ensemble(\n            weighted_avg,\n            sequence,\n            msa_file=msa_file\n        )\n        \n        # Store the 5 structures for this sequence\n        seq_to_structures[target_id] = structures\n    \n    # Fill submission dataframe with predictions\n    print(\"Filling submission dataframe...\")\n    for i, row in tqdm(submission_df.iterrows(), total=len(submission_df), desc=\"Creating submission\"):\n        id_parts = row['ID'].split('_')\n        seq_id = id_parts[0]\n        residue_idx = int(id_parts[1]) - 1\n        \n        if seq_id in seq_to_structures and residue_idx < len(seq_to_structures[seq_id][0]):\n            for struct_idx in range(5):\n                submission_df.at[i, f'x_{struct_idx+1}'] = seq_to_structures[seq_id][struct_idx][residue_idx][0]\n                submission_df.at[i, f'y_{struct_idx+1}'] = seq_to_structures[seq_id][struct_idx][residue_idx][1]\n                submission_df.at[i, f'z_{struct_idx+1}'] = seq_to_structures[seq_id][struct_idx][residue_idx][2]\n    \n    # Save submission file\n    submission_file = os.path.join(OUTPUT_DIR, 'submission.csv')\n    submission_df.to_csv(submission_file, index=False)\n    \n    print(f\"Submission saved to {submission_file}\")\n    print(f\"File size: {os.path.getsize(submission_file) / (1024 * 1024):.2f} MB\")\n    \n    return submission_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:09:09.636528Z","iopub.execute_input":"2025-04-19T16:09:09.636778Z","iopub.status.idle":"2025-04-19T16:09:09.645936Z","shell.execute_reply.started":"2025-04-19T16:09:09.636758Z","shell.execute_reply":"2025-04-19T16:09:09.645190Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===== MAIN EXECUTION =====\ndef main(use_external=True, max_external_structures=50, num_models=5, offline_mode=True):\n    \"\"\"\n    Main execution function.\n    \"\"\"\n    try:\n        # Set offline mode if requested\n        if offline_mode:\n            os.environ['OFFLINE_MODE'] = 'True'\n            print(\"Running in OFFLINE MODE - no internet access will be used\")\n        \n        # Check dependencies\n        if not check_and_install_dependencies():\n            print(\"Please install required dependencies before running\")\n            return None\n            \n        print(\"Starting RNA 3D structure prediction pipeline...\")\n        start_time = time.time()\n        \n        # Load and process data including external structures\n        X_train, y_train, X_valid, y_valid, data = load_all_data(\n            use_external=use_external,\n            max_external_structures=max_external_structures\n        )\n        \n        # Run multi-model ensemble\n        submission_df = run_multi_model_ensemble(\n            X_train, y_train, X_valid, y_valid,\n            data[\"test_sequences\"], data[\"sample_submission\"],\n            num_models=num_models\n        )\n        \n        print(f\"Pipeline completed in {(time.time() - start_time) / 60:.2f} minutes\")\n        return submission_df\n    \n    except Exception as e:\n        print(f\"Error in main execution: {str(e)}\")\n        traceback.print_exc()\n        return None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:09:12.659242Z","iopub.execute_input":"2025-04-19T16:09:12.659512Z","iopub.status.idle":"2025-04-19T16:09:12.665145Z","shell.execute_reply.started":"2025-04-19T16:09:12.659482Z","shell.execute_reply":"2025-04-19T16:09:12.664380Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    # Parse command line arguments\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-19T16:10:04.463531Z","iopub.execute_input":"2025-04-19T16:10:04.464092Z","iopub.status.idle":"2025-04-19T16:28:26.999521Z","shell.execute_reply.started":"2025-04-19T16:10:04.464070Z","shell.execute_reply":"2025-04-19T16:28:26.998858Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}