{"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":"none","dataSources":[{"sourceId":87793,"databundleVersionId":11553390,"sourceType":"competition"}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:42:31.181277Z","iopub.execute_input":"2025-04-20T00:42:31.181653Z","iopub.status.idle":"2025-04-20T00:42:36.178026Z","shell.execute_reply.started":"2025-04-20T00:42:31.181624Z","shell.execute_reply":"2025-04-20T00:42:36.176833Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 1: Data Acquisition and Preprocessing","metadata":{}},{"cell_type":"markdown","source":"### 1.1: Obtain the Dataset:","metadata":{}},{"cell_type":"code","source":"# This is a placeholder - replace with actual download/access code\nimport os\n\ndata_dir = '/kaggle/input/stanford-rna-3d-folding/'\nif not os.path.exists(data_dir):\n    os.makedirs(data_dir)\n    print(f\"Created directory: {data_dir}\")\n    print(\"Please download the Stanford RNA 3D Folding dataset and place it in:\", data_dir)\nelse:\n    print(f\"Data directory exists: {data_dir}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:43:12.776310Z","iopub.execute_input":"2025-04-20T00:43:12.776650Z","iopub.status.idle":"2025-04-20T00:43:12.783429Z","shell.execute_reply.started":"2025-04-20T00:43:12.776624Z","shell.execute_reply":"2025-04-20T00:43:12.782151Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 1.2: Parse the Data:","metadata":{}},{"cell_type":"markdown","source":"### Sequence Extraction:","metadata":{}},{"cell_type":"code","source":"def extract_sequence(filepath):\n    with open(filepath, 'r') as f:\n        lines = f.readlines()\n        sequence = ''.join(lines[1:]).strip().upper().replace('T', 'U')\n    return sequence\n\n# Example usage (assuming a file named 'RNA_001.fasta' exists)\nsequence_file = os.path.join(data_dir, 'RNA_001.fasta')\nif os.path.exists(sequence_file):\n    rna_sequence = extract_sequence(sequence_file)\n    print(f\"Extracted sequence: {rna_sequence}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:43:26.471283Z","iopub.execute_input":"2025-04-20T00:43:26.471658Z","iopub.status.idle":"2025-04-20T00:43:26.480441Z","shell.execute_reply.started":"2025-04-20T00:43:26.471622Z","shell.execute_reply":"2025-04-20T00:43:26.479245Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Structure Representation:","metadata":{}},{"cell_type":"code","source":"def parse_pdb(filepath):\n    coords = {}\n    with open(filepath, 'r') as f:\n        for line in f:\n            if line.startswith(\"ATOM\") and line[17:20].strip() in ['A', 'U', 'G', 'C'] and line[12:16].strip() == \"C3'\":\n                residue_number = int(line[22:26].strip())\n                residue_name = line[17:20].strip()\n                x = float(line[30:38])\n                y = float(line[38:46])\n                z = float(line[46:54])\n                if residue_number not in coords:\n                    coords[residue_number] = {'res_name': residue_name, 'coords': (x, y, z)}\n    # Ensure coordinates are ordered by residue number\n    sorted_coords = [coords[i]['coords'] for i in sorted(coords.keys())]\n    return sorted_coords\n\n# Example usage (assuming a file named 'RNA_001.pdb' exists)\npdb_file = os.path.join(data_dir, 'RNA_001.pdb')\nif os.path.exists(pdb_file):\n    c3_prime_coordinates = parse_pdb(pdb_file)\n    print(f\"Extracted C3' coordinates (first 5): {c3_prime_coordinates[:5]}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:43:33.031209Z","iopub.execute_input":"2025-04-20T00:43:33.031590Z","iopub.status.idle":"2025-04-20T00:43:33.041179Z","shell.execute_reply.started":"2025-04-20T00:43:33.031560Z","shell.execute_reply":"2025-04-20T00:43:33.040145Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Distance Matrix Calculation:","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport torch\n\ndef calculate_distance_matrix(coordinates):\n    n = len(coordinates)\n    dist_matrix = np.zeros((n, n))\n    for i in range(n):\n        for j in range(i + 1, n):\n            dist = np.linalg.norm(np.array(coordinates[i]) - np.array(coordinates[j]))\n            dist_matrix[i, j] = dist\n            dist_matrix[j, i] = dist\n    return torch.tensor(dist_matrix, dtype=torch.float32)\n\nif os.path.exists(pdb_file):\n    c3_prime_coords = parse_pdb(pdb_file)\n    if c3_prime_coords:\n        distance_matrix = calculate_distance_matrix(c3_prime_coords)\n        print(f\"Distance matrix shape: {distance_matrix.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:44:03.896255Z","iopub.execute_input":"2025-04-20T00:44:03.896641Z","iopub.status.idle":"2025-04-20T00:44:09.418383Z","shell.execute_reply.started":"2025-04-20T00:44:03.896610Z","shell.execute_reply":"2025-04-20T00:44:09.417453Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 1.3: Feature Engineering:","metadata":{}},{"cell_type":"code","source":"def one_hot_encode_sequence(sequence):\n    mapping = {'A': 0, 'U': 1, 'G': 2, 'C': 3}\n    encoded_sequence = [mapping[char] for char in sequence]\n    encoded_sequence = torch.nn.functional.one_hot(torch.tensor(encoded_sequence), num_classes=4).float()\n    return encoded_sequence\n\nif os.path.exists(sequence_file):\n    rna_sequence = extract_sequence(sequence_file)\n    encoded_seq = one_hot_encode_sequence(rna_sequence)\n    print(f\"One-hot encoded sequence shape: {encoded_seq.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:44:19.786541Z","iopub.execute_input":"2025-04-20T00:44:19.787578Z","iopub.status.idle":"2025-04-20T00:44:19.795113Z","shell.execute_reply.started":"2025-04-20T00:44:19.787535Z","shell.execute_reply":"2025-04-20T00:44:19.793731Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 1.4: Data Cleaning and Filtering:","metadata":{}},{"cell_type":"markdown","source":"### Handling Missing Files:","metadata":{}},{"cell_type":"code","source":"import os\nfrom torch.utils.data import Dataset\n\nclass RNA3DDataset(Dataset):\n    def __init__(self, data_dir, transform=None):\n        self.data_dir = data_dir\n        self.rna_ids = self._filter_complete_data(data_dir) # Filter in the constructor\n        self.transform = transform\n\n    def _filter_complete_data(self, data_dir):\n        complete_ids = []\n        for f in os.listdir(data_dir):\n            rna_id_base = f.split('.')[0]\n            if f.endswith('.fasta') or f.endswith('.seq'):\n                seq_file = os.path.join(data_dir, f\"{rna_id_base}.{'fasta' if os.path.exists(os.path.join(data_dir, f'{rna_id_base}.fasta')) else 'seq'}\")\n                pdb_file = os.path.join(data_dir, f\"{rna_id_base}.pdb\")\n                if os.path.exists(seq_file) and os.path.exists(pdb_file):\n                    complete_ids.append(rna_id_base)\n        return list(set(complete_ids)) # Ensure unique IDs\n\n    def __len__(self):\n        return len(self.rna_ids)\n\n    def __getitem__(self, idx):\n        rna_id = self.rna_ids[idx]\n        seq_file_ext = 'fasta' if os.path.exists(os.path.join(self.data_dir, f'{rna_id}.fasta')) else 'seq'\n        seq_file = os.path.join(self.data_dir, f\"{rna_id}.{seq_file_ext}\")\n        pdb_file = os.path.join(self.data_dir, f\"{rna_id}.pdb\")\n\n        try:\n            sequence = self._extract_sequence(seq_file)\n            structure_data = self._parse_pdb(pdb_file)\n            c3_prime_coords = structure_data.get('C3\\'')\n            if c3_prime_coords is None or len(sequence) != len(c3_prime_coords):\n                print(f\"Warning: Sequence/structure mismatch or missing C3' for {rna_id}\")\n                return None\n\n            distance_matrix = self._calculate_distance_matrix(c3_prime_coords)\n            sample = {'sequence': sequence, 'distance_matrix': distance_matrix}\n            if self.transform:\n                sample = self.transform(sample)\n            return sample\n        except FileNotFoundError:\n            print(f\"Warning: File not found for {rna_id}\")\n            return None\n\n    def _extract_sequence(self, filepath):\n        with open(filepath, 'r') as f:\n            lines = f.readlines()\n            sequence = ''.join(lines[1:]).strip().upper().replace('T', 'U')\n        return sequence\n\n    def _parse_pdb(self, filepath):\n        coords = {}\n        with open(filepath, 'r') as f:\n            for line in f:\n                if line.startswith(\"ATOM\") and line[17:20].strip() in ['A', 'U', 'G', 'C']:\n                    atom_name = line[12:16].strip()\n                    residue_number = int(line[22:26].strip())\n                    residue_name = line[17:20].strip()\n                    x = float(line[30:38])\n                    y = float(line[38:46])\n                    z = float(line[46:54])\n                    if residue_number not in coords:\n                        coords[residue_number] = {'res_name': residue_name, 'coords': {}}\n                    coords[residue_number]['coords'][atom_name] = (x, y, z)\n        # Reformat to get C3' coordinates in order\n        c3_prime_coords_list = []\n        sorted_residues = sorted(coords.keys())\n        for res_num in sorted_residues:\n            if 'C3\\'' in coords[res_num]['coords']:\n                c3_prime_coords_list.append(coords[res_num]['coords']['C3\\''])\n        return {'C3\\'': c3_prime_coords_list}\n\n    def _calculate_distance_matrix(self, coordinates):\n        n = len(coordinates)\n        dist_matrix = np.zeros((n, n))\n        for i in range(n):\n            for j in range(i + 1, n):\n                dist = np.linalg.norm(np.array(coordinates[i]) - np.array(coordinates[j]))\n                dist_matrix[i, j] = dist\n                dist_matrix[j, i] = dist\n        return torch.tensor(dist_matrix, dtype=torch.float32)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:44:24.951598Z","iopub.execute_input":"2025-04-20T00:44:24.952597Z","iopub.status.idle":"2025-04-20T00:44:24.971352Z","shell.execute_reply.started":"2025-04-20T00:44:24.952545Z","shell.execute_reply":"2025-04-20T00:44:24.970163Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Handling Sequences with Unusual Characters:","metadata":{}},{"cell_type":"code","source":"def _extract_sequence(self, filepath):\n    with open(filepath, 'r') as f:\n        lines = f.readlines()\n        sequence = ''.join(lines[1:]).strip().upper().replace('T', 'U')\n        if not all(char in ['A', 'U', 'G', 'C'] for char in sequence):\n            print(f\"Warning: Sequence contains unusual characters in {filepath}. Skipping.\")\n            return None # Indicate invalid sequence\n        return sequence\n\n# Update __getitem__ to handle None return from _extract_sequence\ndef __getitem__(self, idx):\n    if idx < 0 or idx >= len(self.file_list):\n        raise IndexError(f\"Index {idx} out of range\")\n\n    filename = self.file_list[idx]\n    seq_file = os.path.join(self.data_dir, filename + self.seq_suffix)\n    label_file = os.path.join(self.data_dir, filename + self.label_suffix)\n\n    try:\n        sequence = self._extract_sequence(seq_file)\n        if sequence is None:\n            return None # Skip this item if the sequence is invalid\n\n        with open(label_file, 'r') as f:\n            label = int(f.readline().strip())\n\n        if self.transform:\n            sequence = self.transform(sequence)\n        if self.target_transform:\n            label = self.target_transform(label)\n\n        return sequence, label\n\n    except FileNotFoundError:\n        print(f\"Error: Sequence or label file not found for {filename}. Skipping.\")\n        return None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:44:32.536377Z","iopub.execute_input":"2025-04-20T00:44:32.536770Z","iopub.status.idle":"2025-04-20T00:44:32.545371Z","shell.execute_reply.started":"2025-04-20T00:44:32.536741Z","shell.execute_reply":"2025-04-20T00:44:32.544288Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Handling PDB Files with Inconsistencies:","metadata":{}},{"cell_type":"code","source":"def _parse_pdb(self, filepath):\n    coords = {}\n    models = []\n    current_model = None\n    with open(filepath, 'r') as f:\n        for line in f:\n            if line.startswith(\"MODEL\"):\n                current_model = []\n            elif line.startswith(\"ENDMDL\"):\n                if current_model:\n                    models.append(current_model)\n                current_model = None\n            elif current_model is not None and line.startswith(\"ATOM\") and line[17:20].strip() in ['A', 'U', 'G', 'C']:\n                current_model.append(line)\n            elif current_model is None and line.startswith(\"ATOM\") and line[17:20].strip() in ['A', 'U', 'G', 'C']:\n                models.append([line]) # Handle single model PDBs\n\n    if not models:\n        return {'C3\\'': None} # No valid ATOM records found\n\n    # Use the first model\n    atom_lines = models[0]\n    residue_coords = {}\n    for line in atom_lines:\n        atom_name = line[12:16].strip()\n        residue_number = int(line[22:26].strip())\n        residue_name = line[17:20].strip()\n        x = float(line[30:38])\n        y = float(line[38:46])\n        z = float(line[46:54])\n        if residue_number not in residue_coords:\n            residue_coords[residue_number] = {'res_name': residue_name, 'coords': {}}\n        residue_coords[residue_number]['coords'][atom_name] = (x, y, z)\n\n    c3_prime_coords_list = []\n    sorted_residues = sorted(residue_coords.keys())\n    for res_num in sorted_residues:\n        if 'C3\\'' in residue_coords[res_num]['coords']:\n            c3_prime_coords_list.append(residue_coords[res_num]['coords']['C3\\''])\n        else:\n            print(f\"Warning: Missing C3' atom in residue {res_num} of {filepath}\")\n            return {'C3\\'': None} # Indicate missing C3'\n\n    if len(c3_prime_coords_list) == 0:\n        print(f\"Warning: No C3' coordinates found in {filepath}\")\n        return {'C3\\'': None}\n\n    return {'C3\\'': c3_prime_coords_list}\n\n# Update __getitem__ to handle None return from _parse_pdb\ndef __getitem__(self, idx):\n    rna_id = self.ids[idx]\n    seq_file = os.path.join(self.seq_dir, f\"{rna_id}.fasta\")\n    pdb_file = os.path.join(self.pdb_dir, f\"{rna_id}.pdb\")\n\n    try:\n        sequence = self._extract_sequence(seq_file)\n        if sequence is None:\n            return None\n        structure_data = self._parse_pdb(pdb_file)\n        c3_prime_coords = structure_data.get('C3\\'')\n        if c3_prime_coords is None or len(sequence) != len(c3_prime_coords):\n            print(f\"Warning: Sequence/structure mismatch or missing C3' for {rna_id}\")\n            return None\n\n        # Create a list of coordinate tuples\n        coords_list = [coord for coord in c3_prime_coords]\n        # Convert sequence to numerical representation (if needed)\n        numerical_sequence = [self.letter_to_int[base] for base in sequence]\n\n        if self.transform:\n            coords_list = self.transform(coords_list)\n            numerical_sequence = self.transform(numerical_sequence) # Apply same transform if applicable\n\n        return numerical_sequence, torch.tensor(np.array(coords_list), dtype=torch.float32)\n\n    except FileNotFoundError:\n        print(f\"Warning: Sequence or PDB file not found for {rna_id}\")\n        return None\n    except Exception as e:\n        print(f\"Error processing {rna_id}: {e}\")\n        return None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:44:37.541247Z","iopub.execute_input":"2025-04-20T00:44:37.541570Z","iopub.status.idle":"2025-04-20T00:44:37.556296Z","shell.execute_reply.started":"2025-04-20T00:44:37.541546Z","shell.execute_reply":"2025-04-20T00:44:37.555224Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Filtering Based on Sequence Length or Structure Quality Metrics","metadata":{}},{"cell_type":"code","source":"class RNA3DDataset(Dataset):\n    def __init__(self, data_dir, min_len=10, max_len=500, transform=None):\n        self.data_dir = data_dir\n        self.min_len = min_len\n        self.max_len = max_len\n        self.rna_ids = self._filter_data_by_length(data_dir)\n        self.transform = transform\n\n    def _filter_data_by_length(self, data_dir):\n        valid_ids = []\n        for f in os.listdir(data_dir):\n            rna_id_base = f.split('.')[0]\n            if f.endswith('.fasta') or f.endswith('.seq'):\n                seq_file = os.path.join(data_dir, f\"{rna_id_base}.{'fasta' if os.path.exists(os.path.join(data_dir, f'{rna_id_base}.fasta')) else 'seq'}\")\n                pdb_file = os.path.join(data_dir, f\"{rna_id_base}.pdb\")\n                if os.path.exists(seq_file) and os.path.exists(pdb_file):\n                    sequence = self._extract_sequence_for_length_check(seq_file)\n                    if sequence and self.min_len <= len(sequence) <= self.max_len:\n                        valid_ids.append(rna_id_base)\n        return list(set(valid_ids))\n\n    def _extract_sequence_for_length_check(self, filepath):\n        try:\n            with open(filepath, 'r') as f:\n                lines = f.readlines()\n                sequence = ''.join(lines[1:]).strip().upper().replace('T', 'U')\n                return sequence\n        except Exception:\n            return None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:44:45.016379Z","iopub.execute_input":"2025-04-20T00:44:45.016670Z","iopub.status.idle":"2025-04-20T00:44:45.024977Z","shell.execute_reply.started":"2025-04-20T00:44:45.016645Z","shell.execute_reply":"2025-04-20T00:44:45.024006Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 1.5: Data Splitting:","metadata":{}},{"cell_type":"code","source":"import os\nfrom sklearn.model_selection import train_test_split\nimport random\n\ndata_dir = \"/kaggle/input/stanford-rna-3d-folding/MSA/\"\n\n# Assuming you have a list of RNA identifiers (e.g., filenames without extensions)\nall_rna_ids = [f.split('.')[0] for f in os.listdir(data_dir) if f.endswith('.fasta')]\nrandom.shuffle(all_rna_ids)\n\ntrain_ids, temp_ids = train_test_split(all_rna_ids, test_size=0.3, random_state=42)\nval_ids, test_ids = train_test_split(temp_ids, test_size=0.5, random_state=42)\n\nprint(f\"Number of training samples: {len(train_ids)}\")\nprint(f\"Number of validation samples: {len(val_ids)}\")\nprint(f\"Number of test samples: {len(test_ids)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:44:49.986336Z","iopub.execute_input":"2025-04-20T00:44:49.986627Z","iopub.status.idle":"2025-04-20T00:44:50.934775Z","shell.execute_reply.started":"2025-04-20T00:44:49.986607Z","shell.execute_reply":"2025-04-20T00:44:50.933803Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 2: Model Selection","metadata":{}},{"cell_type":"markdown","source":"### For Distance Matrix Prediction:","metadata":{}},{"cell_type":"code","source":"import torch.nn as nn\n\nclass DistancePredictorCNN(nn.Module):\n    def __init__(self, seq_len, num_filters=32, kernel_size=3):\n        super(DistancePredictorCNN, self).__init__()\n        self.conv1 = nn.Conv1d(4, num_filters, kernel_size, padding='same') # Input channels = 4 (one-hot)\n        self.relu = nn.ReLU()\n        self.conv2 = nn.Conv1d(num_filters, num_filters, kernel_size, padding='same')\n        # ... more convolutional layers ...\n        self.flatten = nn.Flatten()\n        self.fc = nn.Linear(num_filters * seq_len * seq_len, seq_len * seq_len) # Output is flattened distance matrix\n        self.seq_len = seq_len\n\n    def forward(self, x):\n        x = self.relu(self.conv1(x.transpose(1, 2))) # Transpose for Conv1D\n        x = self.relu(self.conv2(x))\n        # ... more convolutional layers ...\n        x = self.flatten(x.unsqueeze(-1).unsqueeze(-1)) # Prepare for FC layer\n        x = self.fc(x)\n        return x.view(self.seq_len, self.seq_len) # Reshape to distance matrix\n\n# Example instantiation\nexample_seq_len = 100\ndistance_model = DistancePredictorCNN(example_seq_len)\nprint(distance_model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:44:57.701905Z","iopub.execute_input":"2025-04-20T00:44:57.702505Z","iopub.status.idle":"2025-04-20T00:45:35.262322Z","shell.execute_reply.started":"2025-04-20T00:44:57.702476Z","shell.execute_reply":"2025-04-20T00:45:35.261329Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### For Torsion Angle Prediction:","metadata":{}},{"cell_type":"code","source":"class TorsionAnglePredictorRNN(nn.Module):\n    def __init__(self, input_size=4, hidden_size=64, num_layers=2, output_size=4): # Example: 4 torsion angles per residue\n        super(TorsionAnglePredictorRNN, self).__init__()\n        self.lstm = nn.LSTM(input_size, hidden_size, num_layers, batch_first=True)\n        self.fc = nn.Linear(hidden_size, output_size)\n\n    def forward(self, x):\n        out, _ = self.lstm(x)\n        out = self.fc(out)\n        return out\n\n# Example instantiation\nexample_seq_len = 100\ntorsion_model = TorsionAnglePredictorRNN(input_size=4, hidden_size=64, output_size=4, num_layers=2)\nprint(torsion_model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:45:46.816332Z","iopub.execute_input":"2025-04-20T00:45:46.816740Z","iopub.status.idle":"2025-04-20T00:45:46.827975Z","shell.execute_reply.started":"2025-04-20T00:45:46.816700Z","shell.execute_reply":"2025-04-20T00:45:46.826950Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 3: Model Development and Training","metadata":{"execution":{"iopub.status.busy":"2025-04-16T07:45:53.611524Z","iopub.execute_input":"2025-04-16T07:45:53.611812Z","iopub.status.idle":"2025-04-16T07:45:53.638144Z","shell.execute_reply.started":"2025-04-16T07:45:53.611793Z","shell.execute_reply":"2025-04-16T07:45:53.636951Z"}}},{"cell_type":"markdown","source":"### Define the Loss Function:","metadata":{}},{"cell_type":"code","source":"distance_criterion = nn.MSELoss()\ntorsion_criterion = nn.MSELoss() # Or other suitable regression loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:45:55.946598Z","iopub.execute_input":"2025-04-20T00:45:55.946966Z","iopub.status.idle":"2025-04-20T00:45:55.952005Z","shell.execute_reply.started":"2025-04-20T00:45:55.946943Z","shell.execute_reply":"2025-04-20T00:45:55.950868Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Choose an Optimizer:","metadata":{}},{"cell_type":"code","source":"import torch.optim as optim\n\ndistance_optimizer = optim.Adam(distance_model.parameters(), lr=0.001)\ntorsion_optimizer = optim.Adam(torsion_model.parameters(), lr=0.001)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:46:00.821030Z","iopub.execute_input":"2025-04-20T00:46:00.821413Z","iopub.status.idle":"2025-04-20T00:46:04.486804Z","shell.execute_reply.started":"2025-04-20T00:46:00.821386Z","shell.execute_reply":"2025-04-20T00:46:04.485507Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Train the Model:","metadata":{}},{"cell_type":"code","source":"def train_distance_model(model, dataloader, criterion, optimizer, num_epochs=10):\n    model.train()\n    for epoch in range(num_epochs):\n        total_loss = 0\n        for batch in dataloader:\n            if batch:\n                sequences = torch.stack([item['sequence'] for item in batch])\n                distance_matrices = torch.stack([item['distance_matrix'] for item in batch])\n\n                optimizer.zero_grad()\n                predictions = model(sequences)\n                loss = criterion(predictions, distance_matrices)\n                loss.backward()\n                optimizer.step()\n                total_loss += loss.item()\n        print(f\"Epoch {epoch+1}, Loss: {total_loss / len(dataloader)}\")\n\n# Example training loop (assuming you have a DataLoader named 'train_loader')\n# train_distance_model(distance_model, train_loader, distance_criterion, distance_optimizer)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:46:11.272239Z","iopub.execute_input":"2025-04-20T00:46:11.272870Z","iopub.status.idle":"2025-04-20T00:46:11.280974Z","shell.execute_reply.started":"2025-04-20T00:46:11.272838Z","shell.execute_reply":"2025-04-20T00:46:11.279828Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 4: Hyperparameter Tuning","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import DataLoader\n\ndef evaluate_distance_model(model, dataloader, criterion):\n    model.eval()\n    total_loss = 0\n    with torch.no_grad():\n        for batch in dataloader:\n            if batch:\n                sequences = torch.stack([item['sequence'] for item in batch])\n                distance_matrices = torch.stack([item['distance_matrix'] for item in batch])\n                predictions = model(sequences)\n                loss = criterion(predictions, distance_matrices)\n                total_loss += loss.item()\n    return total_loss / len(dataloader)\n\n# Example of a simple hyperparameter search (you'd likely use a more systematic approach)\nlearning_rates = [0.001, 0.0005]\nnum_filters_options = [16, 32]\n\nbest_val_loss = float('inf')\nbest_params = None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:46:16.811413Z","iopub.execute_input":"2025-04-20T00:46:16.811797Z","iopub.status.idle":"2025-04-20T00:46:16.819064Z","shell.execute_reply.started":"2025-04-20T00:46:16.811769Z","shell.execute_reply":"2025-04-20T00:46:16.818010Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from typing import List\nclass RNADistanceDataset(Dataset):\n    def __init__(self, sequences: List[str], distance_matrices: List[torch.Tensor]):\n        self.sequences = sequences\n        self.distance_matrices = distance_matrices\n        self.mapping = {'A': 0, 'C': 1, 'G': 2, 'U': 3}\n        self.seq_len = max(len(seq) for seq in sequences)\n\n    def __len__(self):\n        return len(self.sequences)\n\n    def __getitem__(self, idx):\n        sequence = self.sequences[idx]\n        distance_matrix = self.distance_matrices[idx]\n        encoded_sequence = torch.zeros(self.seq_len, 4)\n        for i, base in enumerate(sequence):\n            encoded_sequence[i, self.mapping[base]] = 1\n        return {'sequence': encoded_sequence, 'distance_matrix': distance_matrix}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:46:30.416677Z","iopub.execute_input":"2025-04-20T00:46:30.417073Z","iopub.status.idle":"2025-04-20T00:46:30.425231Z","shell.execute_reply.started":"2025-04-20T00:46:30.417049Z","shell.execute_reply":"2025-04-20T00:46:30.424189Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 5: Model Evaluation","metadata":{"execution":{"iopub.status.busy":"2025-04-16T08:00:30.920450Z","iopub.execute_input":"2025-04-16T08:00:30.921043Z","iopub.status.idle":"2025-04-16T08:00:30.964811Z","shell.execute_reply.started":"2025-04-16T08:00:30.921017Z","shell.execute_reply":"2025-04-16T08:00:30.963427Z"}}},{"cell_type":"code","source":"def calculate_rmsd(predicted_coords, true_coords):\n    # Implementation of RMSD calculation\n    # This requires aligning the predicted and true structures\n    pass # Replace with actual RMSD calculation\n\ndef evaluate_structure_prediction(model, dataloader):\n    model.eval()\n    all_rmsds = []\n    with torch.no_grad():\n        for batch in dataloader:\n            if batch:\n                sequences = [item['sequence'] for item in batch]\n                # Assuming your model predicts coordinates directly or something from which coords can be derived\n                # true_coords_list = [item['coordinates'] for item in batch]\n                predicted_outputs = model(torch.stack(sequences))\n                for i in range(len(batch)):\n                     predicted_coords = [] # Derive coordinates from model output\n                     for box in raw_output:\n                            x_min, y_min, x_max, y_max = box\n                            center_x = (x_min + x_max) / 2\n                            center_y = (y_min + y_max) / 2\n                            predicted_coords.append((center_x, center_y))\n\n                     true_coords = true_coords_list[i]\n                     if predicted_coords is not None and len(predicted_coords) == len(true_coords):\n                         rmsd = calculate_rmsd(predicted_coords, true_coords)\n                         all_rmsds.append(rmsd)\n    if all_rmsds:\n        print(f\"Average RMSD on test set: {np.mean(all_rmsds)}\")\n    else:\n        print(\"No valid predictions for RMSD calculation.\")\n\n# Assuming you have a test DataLoader named 'test_loader'\n# evaluate_structure_prediction(best_distance_model, test_loader) # If your model predicts something from which structure can be derived","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:52:13.896956Z","iopub.execute_input":"2025-04-20T00:52:13.897408Z","iopub.status.idle":"2025-04-20T00:52:13.906229Z","shell.execute_reply.started":"2025-04-20T00:52:13.897381Z","shell.execute_reply":"2025-04-20T00:52:13.905019Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 6: Iteration and Refinement","metadata":{}},{"cell_type":"markdown","source":"### 6.1: Error Analysis","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import mean_squared_error # Example for distance matrix\n\ndef analyze_errors_distance_matrix(model, dataloader, device=\"cpu\", num_samples=5):\n    model.eval()\n    error_examples = []\n    with torch.no_grad():\n        for i, batch in enumerate(dataloader):\n            if i >= num_samples:\n                break\n            if batch:\n                sequences = torch.stack([item['sequence'] for item in batch]).to(device)\n                true_distance_matrices = torch.stack([item['distance_matrix'] for item in batch]).to(device)\n                rna_ids = [item['id'] for item in batch] # Assuming you included IDs in your dataset\n\n                predicted_distance_matrices = model(sequences)\n\n                for j in range(len(batch)):\n                    true_dm = true_distance_matrices[j].cpu().numpy()\n                    pred_dm = predicted_distance_matrices[j].cpu().numpy()\n                    mse = mean_squared_error(true_dm.flatten(), pred_dm.flatten())\n                    error_examples.append({'id': rna_ids[j], 'true': true_dm, 'predicted': pred_dm, 'mse': mse})\n\n    # Sort by error for visualization\n    error_examples.sort(key=lambda x: x['mse'], reverse=True)\n\n    for example in error_examples:\n        print(f\"RNA ID: {example['id']}, MSE: {example['mse']:.4f}\")\n        plt.figure(figsize=(10, 5))\n        plt.subplot(1, 2, 1)\n        plt.imshow(example['true'], cmap='viridis')\n        plt.title('True Distance Matrix')\n        plt.colorbar()\n        plt.subplot(1, 2, 2)\n        plt.imshow(example['predicted'], cmap='viridis')\n        plt.title('Predicted Distance Matrix')\n        plt.colorbar()\n        plt.show()\n\n# Assuming you have a test_loader and your model is 'best_distance_model'\n# and your dataset returns 'id' in each item\n# analyze_errors_distance_matrix(best_distance_model, test_loader, device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:48:08.981783Z","iopub.execute_input":"2025-04-20T00:48:08.982224Z","iopub.status.idle":"2025-04-20T00:48:08.994242Z","shell.execute_reply.started":"2025-04-20T00:48:08.982196Z","shell.execute_reply":"2025-04-20T00:48:08.993071Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"###  6.2: Feature Engineering","metadata":{}},{"cell_type":"markdown","source":"### Install Biopython:","metadata":{}},{"cell_type":"code","source":"pip install biopython","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:48:20.066600Z","iopub.execute_input":"2025-04-20T00:48:20.066983Z","iopub.status.idle":"2025-04-20T00:48:27.517307Z","shell.execute_reply.started":"2025-04-20T00:48:20.066959Z","shell.execute_reply":"2025-04-20T00:48:27.516054Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from Bio import SeqIO\n# Install RNAfold if not already: conda install -c bioconda rnafold\n\ndef predict_secondary_structure(sequence):\n    import subprocess\n    try:\n        process = subprocess.run(['RNAfold', '-i'], input=sequence.encode('utf-8'), capture_output=True, text=True, check=True)\n        output = process.stdout.strip().split('\\n')[1] # Assuming dot-bracket notation is on the second line\n        return output.split()[0]\n    except subprocess.CalledProcessError as e:\n        print(f\"RNAfold error: {e}\")\n        return None\n\nclass EnhancedRNA3DDataset(RNA3DDataset): # Inherit from your previous dataset class\n    def __getitem__(self, idx):\n        sample = super().__getitem__(idx)\n        if sample is None:\n            return None\n\n        sequence = sample['sequence']\n        # Predict secondary structure\n        secondary_structure = predict_secondary_structure(sequence)\n        if secondary_structure:\n            # Encode secondary structure (e.g., one-hot for each position: paired, unpaired)\n            secondary_features = self._encode_secondary_structure(secondary_structure)\n            sample['secondary_structure'] = secondary_features\n        else:\n            sample['secondary_structure'] = torch.zeros(len(sequence), 2) # Example: all unpaired if prediction fails\n\n        # Add other features here if needed\n        return sample\n\n    def _encode_secondary_structure(self, ss):\n        encoding = []\n        for char in ss:\n            if char == '.':\n                encoding.append([1, 0]) # Unpaired\n            elif char in '()[]{}:<>': # Paired (can distinguish different types if needed)\n                encoding.append([0, 1]) # Paired\n            else:\n                encoding.append([0.5, 0.5]) # Unknown\n        return torch.tensor(encoding, dtype=torch.float32)\n\n# Example usage:\n# enhanced_dataset = EnhancedRNA3DDataset(data_dir, transform=YourExistingTransform)\n# enhanced_dataloader = DataLoader(enhanced_dataset, batch_size=batch_size, shuffle=True, collate_fn=lambda x: [item for item in x if item is not None])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:48:31.096894Z","iopub.execute_input":"2025-04-20T00:48:31.097310Z","iopub.status.idle":"2025-04-20T00:48:31.196151Z","shell.execute_reply.started":"2025-04-20T00:48:31.097280Z","shell.execute_reply":"2025-04-20T00:48:31.194905Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 6.3: Model Architecture Changes","metadata":{}},{"cell_type":"code","source":"import torch.nn as nn\n\n# Example: Adding a Bidirectional LSTM layer to a CNN-based model\nclass HybridCNNLSTM(nn.Module):\n    def __init__(self, seq_len, num_filters=32, kernel_size=3, lstm_hidden=64, lstm_layers=2):\n        super(HybridCNNLSTM, self).__init__()\n        self.conv1 = nn.Conv1d(4, num_filters, kernel_size, padding='same')\n        self.relu = nn.ReLU()\n        # ... more CNN layers ...\n        self.lstm = nn.LSTM(num_filters, lstm_hidden, lstm_layers, batch_first=True, bidirectional=True)\n        self.fc = nn.Linear(lstm_hidden * 2 * seq_len, seq_len * seq_len) # Adjust output size\n\n    def forward(self, x):\n        x = self.relu(self.conv1(x.transpose(1, 2)))\n        # ... more CNN layers ...\n        x = x.transpose(1, 2) # Prepare for LSTM (batch, seq, features)\n        out, _ = self.lstm(x)\n        out = out.reshape(out.size(0), -1) # Flatten for FC\n        out = self.fc(out)\n        return out.view(out.size(0), self.seq_len, self.seq_len)\n\n# Example: Using a Transformer Encoder\nclass TransformerDistancePredictor(nn.Module):\n    def __init__(self, seq_len, num_heads=4, num_layers=2, d_model=64):\n        super(TransformerDistancePredictor, self).__init__()\n        self.embedding = nn.Linear(4, d_model) # Project one-hot to d_model\n        self.transformer_encoder_layer = nn.TransformerEncoderLayer(d_model=d_model, nhead=num_heads)\n        self.transformer_encoder = nn.TransformerEncoder(self.transformer_encoder_layer, num_layers=num_layers)\n        self.fc = nn.Linear(d_model * seq_len, seq_len * seq_len)\n\n    def forward(self, x):\n        embedded = self.embedding(x)\n        encoded = self.transformer_encoder(embedded.transpose(0, 1)).transpose(0, 1) # (batch, seq, d_model)\n        flattened = encoded.reshape(encoded.size(0), -1)\n        out = self.fc(flattened)\n        return out.view(out.size(0), x.size(1), x.size(1))\n\n# ... instantiate and train the new models ...","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:48:34.911674Z","iopub.execute_input":"2025-04-20T00:48:34.912072Z","iopub.status.idle":"2025-04-20T00:48:34.923337Z","shell.execute_reply.started":"2025-04-20T00:48:34.912050Z","shell.execute_reply":"2025-04-20T00:48:34.921963Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 6.4: Loss Function Modification","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\n\n# Example: Using a custom loss that penalizes long-range distance errors more\nclass WeightedMSELoss(nn.Module):\n    def __init__(self, alpha=1.0):\n        super(WeightedMSELoss, self).__init__()\n        self.alpha = alpha\n\n    def forward(self, predicted, target):\n        mse = (predicted - target)**2\n        # Create a weight matrix where long-range distances have higher weights\n        weights = torch.ones_like(target)\n        n = target.size(-1)\n        for i in range(n):\n            for j in range(i + 1, n):\n                distance = abs(i - j)\n                if distance > n // 2: # Example threshold for \"long-range\"\n                    weights[:, i, j] *= self.alpha\n                    weights[:, j, i] *= self.alpha\n        return (mse * weights).mean()\n\n# Example: Using a loss that encourages specific contact patterns (requires defining contacts)\nclass ContactMapLoss(nn.Module):\n    def __init__(self, threshold=8.0): # Distance threshold for contact\n        super(ContactMapLoss, self).__init__()\n        self.threshold = threshold\n        self.bce = nn.BCEWithLogitsLoss() # Binary Cross-Entropy for contact prediction\n\n    def forward(self, predicted_distance_matrix):\n        # Convert distance matrix to contact probability (sigmoid or similar)\n        contact_probabilities = torch.sigmoid(-predicted_distance_matrix) # Closer = higher prob\n\n        # Generate \"ground truth\" contact map from true coordinates (if available in the batch)\n        # This part is complex and depends on how your data is structured\n        # true_contact_map = ...\n\n        # if true_contact_map is not None:\n        #     return self.bce(contact_probabilities, true_contact_map.float())\n        # else:\n        return torch.tensor(0.0, requires_grad=True) # Placeholder if no true contacts\n\n# ... replace your loss function in the training loop ...\n# distance_criterion = WeightedMSELoss(alpha=2.0)\n# distance_criterion = ContactMapLoss(threshold=8.0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:48:39.651750Z","iopub.execute_input":"2025-04-20T00:48:39.652121Z","iopub.status.idle":"2025-04-20T00:48:39.661483Z","shell.execute_reply.started":"2025-04-20T00:48:39.652098Z","shell.execute_reply":"2025-04-20T00:48:39.660113Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"###  6.5: Data Augmentation","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport torch\n\nclass CoordinatePerturbation(object):\n    def __init__(self, max_translation=0.1):\n        self.max_translation = max_translation\n\n    def __call__(self, sample):\n        if 'coordinates' in sample:\n            coords = np.array(sample['coordinates'])\n            translation = np.random.uniform(-self.max_translation, self.max_translation, size=coords.shape)\n            perturbed_coords = coords + translation\n            sample['coordinates'] = perturbed_coords.tolist()\n            # Recalculate distance matrix if you are predicting that\n            if 'distance_matrix' in sample:\n                n = len(perturbed_coords)\n                dist_matrix = np.zeros((n, n))\n                for i in range(n):\n                    for j in range(i + 1, n):\n                        dist = np.linalg.norm(perturbed_coords[i] - perturbed_coords[j])\n                        dist_matrix[i, j] = dist\n                        dist_matrix[j, i] = dist\n                sample['distance_matrix'] = torch.tensor(dist_matrix, dtype=torch.float32)\n        return sample\n\n# Apply this transform during training data loading\n# train_dataset = RNA3DDataset(train_data_dir, transform=transforms.Compose([OneHotEncode(), CoordinatePerturbation()]))\n# train_loader = DataLoader(train_dataset, ...)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T00:48:46.136779Z","iopub.execute_input":"2025-04-20T00:48:46.138078Z","iopub.status.idle":"2025-04-20T00:48:46.146813Z","shell.execute_reply.started":"2025-04-20T00:48:46.138031Z","shell.execute_reply":"2025-04-20T00:48:46.145592Z"}},"outputs":[],"execution_count":null}]}