{"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"}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install Bio","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-16T12:58:11.934531Z","iopub.execute_input":"2025-04-16T12:58:11.934783Z","iopub.status.idle":"2025-04-16T12:58:17.529149Z","shell.execute_reply.started":"2025-04-16T12:58:11.934762Z","shell.execute_reply":"2025-04-16T12:58:17.528458Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Active Summary:\n*******************************************************************************************\nThe code implements a comprehensive RNA 3D structure prediction pipeline \nthat integrates deep learning architectures—including convolutional neural networks, \ntransformer encoders with self-attention, and graph convolution layers—with rigorous data preprocessing, normalization, feature augmentation, and evaluation frameworks (using TM-score loss, test-time augmentation, holdout evaluation, and cross-validation) \nto effectively learn and predict nucleotide-based 3D coordinates.\n*******************************************************************************************","metadata":{}},{"cell_type":"markdown","source":"Here is a comprehensive explanation of the attached code. This code implements a complete deep learning pipeline for predicting RNA three-dimensional (3D) structure from RNA sequences. \n\nFurthermore, it spans multiple stages—including data loading and preprocessing, feature engineering, model definition (with both a basic convolutional network and an improved model architecture that incorporates attention, transformer, and graph convolution techniques), training, evaluation, test-time augmentation, and final prediction generation.\n\nBelow, each major section of the code is explained in depth:\n\n# 1. Setup, Imports, and Configuration\nImports\n\nThe code begins with importing a variety of Python libraries that serve different purposes:\n\n    Standard libraries:\n\n        os for file and directory operations.\n\n        gc, time, and random for memory management, timing, and randomness.\n\n    Numerical and Data Processing:\n\n        numpy (as np) for numerical operations.\n\n        pandas (as pd) for handling data frames (CSV reading/writing).\n\n    Deep Learning with PyTorch:\n\n        torch and its submodules (nn, optim, DataLoader, etc.) for building, training, and optimizing neural networks.\n\n    Bioinformatics:\n\n        Bio.SeqIO for reading FASTA files containing multiple sequence alignments (MSA).\n\n    Visualization and Progress Reporting:\n\n        matplotlib.pyplot for plotting training histories and model evaluation visualizations.\n\n        tqdm for progress bars.\n\n    Model Interpretation:\n\n        shap for computing SHAP values (a popular technique for explaining model predictions).\n\n    Scikit-learn:\n\n        Various modules for splitting data, performance metrics (e.g., mean squared error, classification report), scaling features, and cross-validation.\n\nConfiguration Parameters\n\nAfter the imports, some configuration variables are defined:\n\n    DATA_PATH: Path to the dataset directory.\n\n    SEQ_LENGTH: Maximum sequence length used for modeling (fixed at 256).\n\n    BATCH_SIZE: The number of samples per mini-batch during training.\n\n    EPOCHS: Maximum number of epochs for training.\n\n    DEVICE: This setting selects CUDA (GPU) if available; otherwise, it falls back to CPU.\n\n    NUM_PREDICTIONS: Number of structure predictions made per sample (used later for multi-prediction settings).\n\nSeeding for Reproducibility\n\nThe set_seed() function is defined and then called. It sets the seed for the Python built-in random module, NumPy, and various PyTorch modules (both CPU and GPU) to ensure that experiments are reproducible. It also ensures deterministic behavior with respect to CUDA operations and sets PYTHONHASHSEED to maintain hash consistency.\nMemory Optimization and Cleanup\n\nTwo helper functions:\n\n    reduce_mem_usage(df):\n    Iterates over the columns of a pandas DataFrame and downcasts numerical columns to lower precision data types when possible, which reduces memory usage. It prints the memory usage before and after optimization.\n\n    clear_memory():\n    This function forces garbage collection and clears PyTorch’s CUDA cache in order to free up GPU memory.\n\n# 2. Data Loading, Preprocessing, and Feature EngineeringLoading Data\n\n    load_data():\n    Loads several CSV files from the provided DATA_PATH. These CSVs include training sequences, validation sequences, test sequences, and associated labels. It also loads a sample submission file for later inference.\n\nData Verification\n\n    verify_data_splits(train_df, val_df, test_df):\n    Checks for overlapping target IDs among the splits (train, validation, and test) to avoid data leakage during training and evaluation.\n\nFeature Engineering Functions\n\n    create_advanced_lag_features(df, window_sizes=[3, 5]):\n    For numeric columns in a DataFrame, computes rolling window statistics (mean, standard deviation, min, and max) over specified window sizes. These advanced features can capture local trends in the data.\n\n    augment_rna_sequence(seq, prob=0.1):\n    Applies random augmentation to RNA sequences. It randomly substitutes characters (representing RNA bases A, C, G, U) with a given probability (default 10%), which can help generalize the model by artificially increasing data variability.\n\n    augment_features(features, noise_std=0.01):\n    Adds Gaussian noise to the input features (using PyTorch operations) to further augment the data during training.\n\nThe Custom Dataset: RNADataset\n\nThis is a subclass of PyTorch’s Dataset that handles both the input features (RNA sequences plus additional MSA-based conservation features) and labels (3D structure coordinates).\n\n    Initialization (__init__):\n    It receives as input:\n\n        A DataFrame with sequences.\n\n        A DataFrame with labels.\n\n        The directory path for MSA files.\n\n        A maximum sequence length (max_len, set from SEQ_LENGTH).\n\n        Flags for whether to perform augmentation and normalization.\n\n    During initialization, it pre-processes each sample and stores the processed data in memory.\n\n    Preprocessing of Sequences (_preprocess_sequence):\n    Maps nucleotide bases (‘A’, ‘C’, ‘G’, ‘U’) into a one-hot encoded representation in a fixed-size array (of shape max_len x 4). Only the first max_len bases of the sequence are used.\n\n    Extraction of MSA Features (_get_msa_features):\n    Loads MSA data from a FASTA file (if present) corresponding to the target ID. It computes:\n\n        The counts of each base in the aligned sequences.\n\n        Frequencies, and derives an entropy measure.\n\n        Conservation score which is then combined with an extra feature (scaled sequence depth).\n\n    If the file does not exist or an error occurs, it returns a zero matrix.\n\n    Label Preprocessing (_preprocess_labels):\n    Retrieves coordinate information from the label DataFrame for a given target ID and processes them:\n\n        Iterates through expected coordinate columns (x_i, y_i, z_i).\n\n        Pads the structure coordinates to match max_len if there are fewer residues or trims if too many.\n\n        The final array has a shape reflecting multiple predictions (set by NUM_PREDICTIONS), and the data is normalized across the coordinates.\n\n    Dataset Creation (_preprocess_data):\n    It loops over the sequence DataFrame:\n\n        Normalization Step (if requested):\n        First, it collects features across all samples to fit a StandardScaler (from scikit-learn) on the flattened feature array.\n\n        Processing and Data Augmentation:\n        For every row, it processes the sequence (with potential augmentation), combines one-hot and MSA features, applies scaling, processes labels, and if augmentation is enabled, adds noise to the coordinate labels.\n\n        The processed sample is a tuple of tensors:\n\n            The features tensor has shape (channels, seq_len) (channels include 4 for one-hot and 2 for conservation).\n\n            The labels tensor is shaped (NUM_PREDICTIONS, max_len, 3), representing 3D coordinates.\n\n    __len__ and __getitem__ Methods:\n    Define dataset length and indexing, which simply return the preprocessed data.\n\nCustom Collate Function\n\n    custom_collate_fn(batch):\n    This function is used by the DataLoader to combine a list of dataset samples into a batch. It stacks the feature and label tensors and ensures that:\n\n        The features are padded or truncated to a fixed sequence length.\n\n        Similarly, label tensors are padded or truncated to maintain consistent dimensions across the batch.\n\n# 3. Neural Network Architecture Basic Model: RNA3DModel\n\n    Structure:\n\n        Input Channels: 6 (four from one-hot encoding and two from conservation features).\n\n        Convolutional Blocks:\n        Uses successive 1D convolutional layers with ReLU activations and Batch Normalization.\n\n        Output Layer:\n        A final convolutional layer followed by a Tanh activation, reshapes the output to format the predicted 3D coordinates. The output is rearranged to have a shape where each prediction for the structure is represented as a set of (x, y, z) coordinates across positions.\n\nImproved Model: ImprovedRNA3DModel\n\nThis model introduces several enhancements:\n\n    Initial Convolutional Block:\n    Similar to the basic model but includes dropout layers to regularize the training.\n\n    Self-Attention:\n    A multihead self-attention mechanism (using nn.MultiheadAttention) is applied after the initial convolution. This helps capture long-range dependencies in the sequence.\n\n    Transformer Encoder:\n    A transformer encoder (with configurable layers and heads) refines the sequence representation.\n\n    Graph Convolution Layers:\n    Two custom graph convolution layers (implemented in GraphConvLayer) are then applied. These layers aggregate local neighborhood information by using a sliding window (averaging over adjacent positions) followed by a linear transformation and dropout.\n\n    Post-Processing Block:\n    Once the features are enriched via attention and graph convolutions, a post-conv block refines them before generating the final predictions.\n\n    Output and Confidence Head:\n\n        Output:\n        Similar to the basic model, a convolutional layer produces the structure prediction.\n\n        Confidence Head:\n        An additional head computes confidence scores for the predictions (using a small network that outputs a sigmoid activation). This score can then be used to weight predictions during test-time augmentation (TTA).\n\n    Weight Initialization:\n    There is an _initialize_weights() method to properly initialize layers (using Kaiming and Xavier initialization methods) for better convergence.\n\nTest-Time Augmentation (TTA)\n\n    tta_predict():\n    This function performs Test-Time Augmentation by making multiple forward passes with slight variations in dropout rates and input noise. It:\n\n        Runs multiple iterations over the same input features.\n\n        Each iteration applies a different level of dropout and added noise.\n\n        The predictions are combined by weighting them using softmax-normalized confidence scores.\n\n        Returns a weighted average of the predictions along with average confidence.\n\n# 4. Training Components and Loss Function\nCustom Loss: tm_score_loss\n\n    Purpose:\n    Designed to mimic the TM-score metric—a measure of structural similarity in protein/RNA 3D structure predictions.\n\n    Details:\n\n        The loss computes pairwise Euclidean distances between predicted and true coordinates.\n\n        Uses an adaptive distance threshold (d0) that depends on the sequence length.\n\n        The TM-score components are aggregated and then the negative mean is used as loss.\n\n        Additionally, an L2 regularization term is applied to the predicted structure to avoid extreme values.\n\nAccuracy Computation: compute_accuracy\n\n    Method:\n    Iterates over multiple predictions and computes the fraction of coordinates for which the Euclidean distance between prediction and ground truth falls below a specified threshold.\n\n    Outcome:\n    Returns the best accuracy among the multiple structure predictions.\n\n# 5. Training, Evaluation, and Experiment Management Training and Validation Loop: train_and_validate\n\n    Training Loop:\n\n        Uses the Adam optimizer and a learning rate scheduler (ReduceLROnPlateau) that reduces the learning rate when the validation loss plateaus.\n\n        For each training batch, features are optionally augmented via Gaussian noise, predictions are computed, and the custom TM-score loss is calculated.\n\n        Gradients are computed and clipped (to prevent exploding gradients), and weights are updated.\n\n    Validation Loop:\n\n        After training on all batches in an epoch, the model is evaluated on the validation set.\n\n        The loss and accuracy are computed for monitoring performance.\n\n    Early Stopping and Model Saving:\n\n        The code checks if the validation loss has improved; if so, it saves the model checkpoint and resets the counter for “no improvement.”\n\n        If the model does not improve for a specified number of epochs (patience), training stops early.\n\n    Visualization:\n\n        Loss history and accuracy history are plotted and saved as PNG images.\n\nTest Prediction: test_prediction\n\n    Inference:\n\n        Loads the best model checkpoint.\n\n        Iterates over the test dataset and computes the structure predictions.\n\n        Extracts the primary prediction (typically the first structure out of multiple predictions).\n\n    Submission Generation:\n\n        Formats predictions into a DataFrame by mapping each target ID to its predicted 3D coordinates.\n\n        This output DataFrame is ready to be saved as a submission CSV file.\n\nModel Performance Analysis: analyze_model_performance\n\n    Confidence Histogram:\n    Plots a histogram of confidence scores generated by the model.\n\n    3D Structure Comparison:\n\n        Compares the predicted 3D structure against the ground truth in a 3D scatter plot, connecting sequential residues to visualize the structure.\n\n    SHAP Analysis:\n\n        Attempts to compute SHAP values to explain feature importance for the model’s predictions.\n\n        If SHAP fails (e.g., due to computational constraints), an error is caught and printed.\n\nHoldout Evaluation and Cross-Validation\n\n    evaluate_holdout():\n    Splits a portion of the training data to serve as a holdout set. It then evaluates the model on this unseen data and prints the average loss and accuracy.\n\n    cross_validation():\n    Implements k-fold (StratifiedKFold) cross-validation:\n\n        Uses stratification based on sequence lengths.\n\n        For each fold, the training and validation sets are created, and the model is reinitialized and trained.\n\n        The average validation metrics across folds are computed.\n\n# 6. Main Execution Pipeline The main() function ties together all the pieces:\n\n    Environment Setup:\n\n        Prints which device (CPU or GPU) is used.\n\n        Loads data using load_data() and checks for overlapping target IDs with verify_data_splits().\n\n    Dataset and DataLoader Initialization:\n\n        Creates instances of the RNADataset for training, validation, and test sets.\n\n        Wraps these datasets in PyTorch DataLoaders with the custom collate function for fixed-length sequences.\n\n    Model Initialization:\n\n        An instance of the improved model (ImprovedRNA3DModel) is created and sent to the appropriate device.\n\n    Training:\n\n        Calls train_and_validate() to start the training process, while monitoring and saving the best model checkpoint.\n\n    Holdout Evaluation and Visualization:\n\n        Evaluates the model on a holdout set.\n\n        Plots overall loss and accuracy trends including holdout performance.\n\n    Model Analysis:\n\n        Runs further analysis (confidence distributions, SHAP feature importance, etc.) using analyze_model_performance().\n\n    Test-Time Augmentation and Predictions:\n\n        Generates predictions with TTA via tta_predict().\n\n        Formats predictions to create a final submission file using test_prediction().\n\n    Optional Cross-Validation:\n\n        Provides an optional block (currently disabled) for running cross-validation to further assess model stability.\n\n    Resource Cleanup:\n\n        Clears memory (garbage collection and CUDA cache clearing) after execution to free up resources.\n\n    Execution Trigger:\n\n        The script runs main() when executed as the main module.\n\n","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom Bio import SeqIO\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nimport shap\nfrom sklearn.model_selection import train_test_split, KFold, StratifiedKFold\nfrom sklearn.metrics import mean_squared_error, classification_report\nfrom sklearn.preprocessing import StandardScaler\nimport gc\nimport random\nimport time\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nfrom torch.nn import functional as F\n\n# =============================================================================\n# 1. Configuration, Seeds, and Memory Management\n# =============================================================================\n\nDATA_PATH = \"/kaggle/input/stanford-rna-3d-folding/\"\nSEQ_LENGTH = 256\nBATCH_SIZE = 32\nEPOCHS = 50\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nNUM_PREDICTIONS = 5  # Number of structure predictions\n\ndef set_seed(seed=42):\n    \"\"\"Set seeds for reproducibility across all random modules.\"\"\"\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    print(f\"Seeds set to {seed} for reproducibility\")\n\nset_seed()\n\ndef reduce_mem_usage(df):\n    \"\"\"Downcast numeric columns to reduce memory usage.\"\"\"\n    start_mem = df.memory_usage().sum() / 1024**2\n    print(f\"Memory usage of dataframe is {start_mem:.2f} MB\")\n    for col in df.columns:\n        if df[col].dtype != object:\n            c_min = df[col].min()\n            c_max = df[col].max()\n            if str(df[col].dtype)[:3] == \"int\":\n                if c_min > np.iinfo(np.int8).min and c_max < np.iinfo(np.int8).max:\n                    df[col] = df[col].astype(np.int8)\n                elif c_min > np.iinfo(np.int16).min and c_max < np.iinfo(np.int16).max:\n                    df[col] = df[col].astype(np.int16)\n                elif c_min > np.iinfo(np.int32).min and c_max < np.iinfo(np.int32).max:\n                    df[col] = df[col].astype(np.int32)\n                else:\n                    df[col] = df[col].astype(np.int64)\n            else:\n                df[col] = df[col].astype(np.float32)\n    end_mem = df.memory_usage().sum() / 1024**2\n    print(f\"Memory usage after optimization: {end_mem:.2f} MB; decreased by {100 * (start_mem - end_mem) / start_mem:.1f}%\")\n    return df\n\ndef clear_memory():\n    \"\"\"Force garbage collection and clear CUDA cache.\"\"\"\n    gc.collect()\n    torch.cuda.empty_cache()\n\n# =============================================================================\n# 2. Data Loading, Preprocessing, and Feature Engineering\n# =============================================================================\n\ndef load_data():\n    \"\"\"Load CSV files with sequences, labels, and sample submission.\"\"\"\n    print(\"Loading data...\")\n    train_seqs = pd.read_csv(os.path.join(DATA_PATH, 'train_sequences.csv'))\n    val_seqs = pd.read_csv(os.path.join(DATA_PATH, 'validation_sequences.csv'))\n    test_seqs = pd.read_csv(os.path.join(DATA_PATH, 'test_sequences.csv'))\n    train_labels = pd.read_csv(os.path.join(DATA_PATH, 'train_labels.csv'))\n    val_labels = pd.read_csv(os.path.join(DATA_PATH, 'validation_labels.csv'))\n    sample_submission = pd.read_csv(os.path.join(DATA_PATH, 'sample_submission.csv'))\n    return train_seqs, val_seqs, test_seqs, train_labels, val_labels, sample_submission\n\ndef verify_data_splits(train_df, val_df, test_df):\n    \"\"\"Ensure that there is no data leakage between splits.\"\"\"\n    train_ids = set(train_df['target_id'])\n    val_ids = set(val_df['target_id'])\n    test_ids = set(test_df['target_id'])\n    if not (train_ids.isdisjoint(val_ids) and train_ids.isdisjoint(test_ids) and val_ids.isdisjoint(test_ids)):\n        print(\"WARNING: Overlapping target IDs detected among splits.\")\n    else:\n        print(\"Data splits verified: No overlapping target IDs.\")\n\ndef create_advanced_lag_features(df, window_sizes=[3, 5]):\n    \"\"\"Add rolling statistics as lag features.\"\"\"\n    for col in df.select_dtypes(include=[np.number]).columns:\n        for window in window_sizes:\n            df[f'{col}_roll_mean_{window}'] = df[col].rolling(window=window, min_periods=1).mean()\n            df[f'{col}_roll_std_{window}'] = df[col].rolling(window=window, min_periods=1).std()\n            df[f'{col}_roll_min_{window}'] = df[col].rolling(window=window, min_periods=1).min()\n            df[f'{col}_roll_max_{window}'] = df[col].rolling(window=window, min_periods=1).max()\n    return df\n\n# Data augmentation for RNA sequences\ndef augment_rna_sequence(seq, prob=0.1):\n    \"\"\"Randomly substitute bases in RNA sequence.\"\"\"\n    bases = ['A', 'C', 'G', 'U']\n    seq_chars = list(seq)\n    for i in range(len(seq_chars)):\n        if random.random() < prob:\n            seq_chars[i] = random.choice([b for b in bases if b != seq_chars[i]])\n    return ''.join(seq_chars)\n\n# Function to add Gaussian noise to features\ndef augment_features(features, noise_std=0.01):\n    noise = torch.randn_like(features) * noise_std\n    return features + noise\n\nclass RNADataset(Dataset):\n    def __init__(self, seq_df, label_df, msa_dir, max_len=SEQ_LENGTH, augment=False, normalize=True):\n        self.seq_df = seq_df\n        self.label_df = label_df\n        self.msa_dir = msa_dir\n        self.max_len = max_len\n        self.augment = augment\n        self.normalize = normalize\n        self.data = []\n        self.feature_scaler = StandardScaler()\n        self._preprocess_data()\n        \n    def _preprocess_sequence(self, seq):\n        mapping = {'A': 0, 'C': 1, 'G': 2, 'U': 3}\n        one_hot = np.zeros((self.max_len, 4), dtype=np.float32)\n        for i, base in enumerate(seq[:self.max_len]):\n            if base in mapping:\n                one_hot[i, mapping[base]] = 1.0\n        return one_hot\n    \n    def _get_msa_features(self, target_id):\n        msa_path = os.path.join(self.msa_dir, f\"{target_id}.MSA.fasta\")\n        # Return zeros with 2 channels to match expected dimension.\n        if not os.path.exists(msa_path):\n            return np.zeros((self.max_len, 2), dtype=np.float32)\n        try:\n            sequences = [str(rec.seq) for rec in SeqIO.parse(msa_path, 'fasta')]\n            if sequences:\n                counts = np.zeros((self.max_len, 4), dtype=np.float32)\n                for seq in sequences:\n                    for i, c in enumerate(seq[:self.max_len]):\n                        if c == 'A': \n                            counts[i, 0] += 1\n                        elif c == 'C': \n                            counts[i, 1] += 1\n                        elif c == 'G': \n                            counts[i, 2] += 1\n                        elif c == 'U': \n                            counts[i, 1] += 1\n                counts += 1e-5\n                freqs = counts / counts.sum(axis=1, keepdims=True)\n                entropy = -np.sum(freqs * np.log(freqs + 1e-10), axis=1)\n                conservation = 1 - entropy / np.log(4)\n                seq_depth = len(sequences)\n                extra_feature = np.ones(self.max_len, dtype=np.float32) * seq_depth / 100\n                conservation = np.column_stack((conservation, extra_feature))\n                if conservation.shape[0] > self.max_len:\n                    conservation = conservation[:self.max_len, :]\n                elif conservation.shape[0] < self.max_len:\n                    pad_rows = self.max_len - conservation.shape[0]\n                    padding = np.zeros((pad_rows, conservation.shape[1]), dtype=conservation.dtype)\n                    conservation = np.vstack((conservation, padding))\n                return conservation\n            else:\n                return np.zeros((self.max_len, 2), dtype=np.float32)\n        except Exception as e:\n            print(f\"Error processing MSA for {target_id}: {e}\")\n            return np.zeros((self.max_len, 2), dtype=np.float32)\n    \n    def _preprocess_labels(self, target_id):\n        if self.label_df.empty or 'ID' not in self.label_df.columns:\n            return np.zeros((NUM_PREDICTIONS, self.max_len, 3), dtype=np.float32)\n        target_labels = self.label_df[self.label_df['ID'].str.startswith(target_id)]\n        coords_list = []\n        for _, row in target_labels.iterrows():\n            struct_coords = []\n            for i in range(1, 100):\n                x = row.get(f'x_{i}', np.nan)\n                y = row.get(f'y_{i}', np.nan)\n                z = row.get(f'z_{i}', np.nan)\n                if np.isnan(x) or np.isnan(y) or np.isnan(z):\n                    break\n                struct_coords.append([x, y, z])\n            if len(struct_coords) > self.max_len:\n                struct_coords = struct_coords[:self.max_len]\n            elif len(struct_coords) < self.max_len:\n                pad_len = self.max_len - len(struct_coords)\n                struct_coords.extend([[0.0, 0.0, 0.0]] * pad_len)\n            coords_list.append(struct_coords)\n        if len(coords_list) < NUM_PREDICTIONS:\n            for _ in range(NUM_PREDICTIONS - len(coords_list)):\n                coords_list.append([[0.0, 0.0, 0.0]] * self.max_len)\n        else:\n            coords_list = coords_list[:NUM_PREDICTIONS]\n        coords_array = np.array(coords_list, dtype=np.float32)\n        coords_array = (coords_array - np.mean(coords_array)) / (np.std(coords_array) + 1e-8)\n        return coords_array\n    \n    def _preprocess_data(self):\n        all_features = []\n        if self.normalize:\n            print(\"Collecting features for normalization...\")\n            for _, row in tqdm(self.seq_df.iterrows(), total=len(self.seq_df)):\n                target_id = row['target_id']\n                seq = row['sequence']\n                one_hot = self._preprocess_sequence(seq)\n                conservation = self._get_msa_features(target_id)\n                if isinstance(conservation, np.ndarray) and conservation.ndim == 1:\n                    conservation = conservation.reshape(-1, 1)\n                if one_hot.shape[0] != conservation.shape[0]:\n                    conservation = conservation[:one_hot.shape[0], :]\n                features = np.concatenate([one_hot, conservation], axis=1)\n                all_features.append(features.reshape(1, -1)[0])\n            all_features = np.array(all_features)\n            self.feature_scaler.fit(all_features)\n        \n        print(\"Creating dataset...\")\n        for _, row in tqdm(self.seq_df.iterrows(), total=len(self.seq_df)):\n            target_id = row['target_id']\n            seq = row['sequence']\n            if self.augment and random.random() < 0.3:\n                seq = augment_rna_sequence(seq, prob=0.05)\n            one_hot = self._preprocess_sequence(seq)\n            conservation = self._get_msa_features(target_id)\n            if isinstance(conservation, np.ndarray) and conservation.ndim == 1:\n                conservation = conservation.reshape(-1, 1)\n            if one_hot.shape[0] != conservation.shape[0]:\n                conservation = conservation[:one_hot.shape[0], :]\n            features = np.concatenate([one_hot, conservation], axis=1)\n            if self.normalize:\n                flat_features = features.reshape(1, -1)\n                features = self.feature_scaler.transform(flat_features).reshape(features.shape)\n            coords = self._preprocess_labels(target_id)\n            if self.augment and not np.all(coords == 0):\n                noise_scale = 0.01\n                coords += np.random.normal(0, noise_scale, coords.shape)\n            self.data.append((\n                torch.tensor(features.T, dtype=torch.float32),  # (channels, seq_len)\n                torch.tensor(coords, dtype=torch.float32)         # (NUM_PREDICTIONS, max_len, 3)\n            ))\n    \n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self, idx):\n        return self.data[idx]\n\ndef custom_collate_fn(batch):\n    features_list, labels_list = zip(*batch)\n    features = torch.stack(features_list, dim=0)  # (batch, channels, seq_len)\n    fixed_seq_len = SEQ_LENGTH\n    if features.shape[2] < fixed_seq_len:\n        pad_size = fixed_seq_len - features.shape[2]\n        features = F.pad(features, (0, pad_size), \"constant\", 0)\n    elif features.shape[2] > fixed_seq_len:\n        features = features[:, :, :fixed_seq_len]\n    padded_labels = []\n    for label in labels_list:\n        if label.shape[1] < fixed_seq_len:\n            pad = torch.zeros((label.shape[0], fixed_seq_len - label.shape[1], label.shape[2]), dtype=label.dtype)\n            padded_label = torch.cat([label, pad], dim=1)\n        elif label.shape[1] > fixed_seq_len:\n            padded_label = label[:, :fixed_seq_len, :]\n        else:\n            padded_label = label\n        padded_labels.append(padded_label)\n    labels = torch.stack(padded_labels, dim=0)\n    return features, labels\n\n# =============================================================================\n# 3. Neural Network Architecture (Original)\n# =============================================================================\n# Input channels: 6 (4 for one-hot, 2 for conservation)\nclass RNA3DModel(nn.Module):\n    def __init__(self, input_channels=6, seq_length=SEQ_LENGTH, num_structures=NUM_PREDICTIONS):\n        super().__init__()\n        self.num_structures = num_structures\n        self.seq_length = seq_length\n        self.conv_block = nn.Sequential(\n            nn.Conv1d(input_channels, 128, 5, padding=2),\n            nn.ReLU(),\n            nn.BatchNorm1d(128),\n            nn.Conv1d(128, 256, 3, padding=1),\n            nn.ReLU(),\n            nn.BatchNorm1d(256),\n            nn.Conv1d(256, 512, 3, padding=1),\n            nn.ReLU()\n        )\n        self.output = nn.Sequential(\n            nn.Conv1d(512, num_structures * 3, 1),\n            nn.Tanh()\n        )\n    \n    def forward(self, x):\n        x = self.conv_block(x)\n        x = self.output(x)\n        x = x.permute(0, 2, 1)\n        batch_size = x.size(0)\n        x = x.view(batch_size, x.size(1), self.num_structures, 3)\n        return x\n\n# =============================================================================\n# 3.5 Enhanced Model: Attention, Graph-Based, Transformer, Confidence, and TTA\n# =============================================================================\nclass GraphConvLayer(nn.Module):\n    def __init__(self, in_features, out_features, dropout=0.2):\n        super().__init__()\n        self.linear = nn.Linear(in_features, out_features)\n        self.dropout = nn.Dropout(dropout)\n    \n    def forward(self, x):\n        batch, seq_len, feat = x.size()\n        padded = torch.zeros(batch, seq_len + 2, feat, device=x.device)\n        padded[:, 1:-1, :] = x\n        agg = (padded[:, :-2, :] + padded[:, 1:-1, :] + padded[:, 2:, :]) / 3.0\n        out = self.linear(agg)\n        out = self.dropout(out)\n        return out\n\nclass ImprovedRNA3DModel(nn.Module):\n    def __init__(self, input_channels=6, seq_length=SEQ_LENGTH, num_structures=NUM_PREDICTIONS, \n                 nhead=4, num_transformer_layers=2, dropout_rate=0.3):\n        super().__init__()\n        self.num_structures = num_structures\n        self.seq_length = seq_length\n        self.dropout_rate = dropout_rate\n        \n        self.conv_block = nn.Sequential(\n            nn.Conv1d(input_channels, 128, 5, padding=2),\n            nn.ReLU(),\n            nn.BatchNorm1d(128),\n            nn.Dropout(dropout_rate),\n            nn.Conv1d(128, 256, 3, padding=1),\n            nn.ReLU(),\n            nn.BatchNorm1d(256),\n            nn.Dropout(dropout_rate/2)\n        )\n        \n        self.self_attn = nn.MultiheadAttention(embed_dim=256, num_heads=nhead, dropout=dropout_rate, batch_first=True)\n        \n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=256,\n            nhead=nhead,\n            dropout=dropout_rate,\n            dim_feedforward=512,\n            batch_first=True\n        )\n        self.transformer_encoder = nn.TransformerEncoder(encoder_layer, num_layers=num_transformer_layers)\n        \n        self.graph_conv1 = GraphConvLayer(256, 256, dropout=dropout_rate)\n        self.graph_conv2 = GraphConvLayer(256, 256, dropout=dropout_rate/2)\n        \n        self.layer_norm1 = nn.LayerNorm(256)\n        self.layer_norm2 = nn.LayerNorm(256)\n        \n        self.post_block = nn.Sequential(\n            nn.Conv1d(256, 512, 3, padding=1),\n            nn.ReLU(),\n            nn.BatchNorm1d(512),\n            nn.Dropout(dropout_rate/3)\n        )\n        \n        self.out_conv = nn.Sequential(\n            nn.Conv1d(512, num_structures * 3, 1),\n            nn.Tanh()\n        )\n        \n        self.confidence_head = nn.Sequential(\n            nn.Conv1d(512, 64, 1),\n            nn.ReLU(),\n            nn.BatchNorm1d(64),\n            nn.Conv1d(64, 1, 1),\n            nn.Sigmoid()\n        )\n        \n        self._initialize_weights()\n    \n    def _initialize_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv1d):\n                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n                if m.bias is not None:\n                    nn.init.constant_(m.bias, 0)\n            elif isinstance(m, nn.BatchNorm1d):\n                nn.init.constant_(m.weight, 1)\n                nn.init.constant_(m.bias, 0)\n            elif isinstance(m, nn.Linear):\n                nn.init.xavier_normal_(m.weight)\n                if m.bias is not None:\n                    nn.init.constant_(m.bias, 0)\n    \n    def forward(self, x, p_dropout=None):\n        dropout_rate = self.dropout_rate if p_dropout is None else p_dropout\n        x = self.conv_block(x)\n        x = x.transpose(1, 2)  # (batch, seq_len, features)\n        x = F.dropout(x, p=dropout_rate, training=self.training)\n        attn_out, _ = self.self_attn(x, x, x)\n        x = x + F.dropout(attn_out, p=dropout_rate, training=self.training)\n        x = self.layer_norm1(x)\n        x = self.transformer_encoder(x)\n        x = self.layer_norm2(x)\n        gc_out1 = self.graph_conv1(x)\n        x = x + gc_out1\n        gc_out2 = self.graph_conv2(x)\n        x = x + gc_out2\n        x_post = x.transpose(1, 2)  # (batch, features, seq_len)\n        x_post = self.post_block(x_post)\n        out = self.out_conv(x_post)\n        out = out.permute(0, 2, 1)\n        batch_size = out.size(0)\n        structure_pred = out.view(batch_size, out.size(1), self.num_structures, 3)\n        confidence = self.confidence_head(x_post)\n        confidence = confidence.squeeze(1)  # (batch, seq_len)\n        return structure_pred, confidence\n\ndef tta_predict(model, loader, tta_iterations=5):\n    \"\"\"\n    Test-Time Augmentation: perform multiple predictions with slight random noise added,\n    and vary dropout rate. The predictions are then weighted by softmax-normalized confidence.\n    \"\"\"\n    model.eval()\n    all_preds = []\n    all_confidences = []\n    \n    with torch.no_grad():\n        for features, _ in tqdm(loader, desc=\"Generating TTA predictions\"):\n            batch_preds = []\n            batch_confs = []\n            for tta_iter in range(tta_iterations):\n                p_dropout = 0.1 + (tta_iter * 0.05)\n                noise_level = 0.005 + (tta_iter * 0.003)\n                noise = torch.randn_like(features) * noise_level\n                aug_features = features + noise\n                aug_features = aug_features.to(DEVICE)\n                structure_pred, confidence = model(aug_features, p_dropout=p_dropout)\n                batch_preds.append(structure_pred.cpu())\n                batch_confs.append(confidence.cpu())\n            stacked_preds = torch.stack(batch_preds, dim=0)  # [TTA, batch, seq_len, num_structures, 3]\n            stacked_conf = torch.stack(batch_confs, dim=0)     # [TTA, batch, seq_len]\n            norm_conf = F.softmax(stacked_conf, dim=0)  # Normalize over TTA iterations\n            weights = norm_conf.unsqueeze(-1).unsqueeze(-1)  # [TTA, batch, seq_len, 1, 1]\n            weighted_pred = (stacked_preds * weights).sum(dim=0)\n            avg_conf = stacked_conf.mean(dim=0)\n            all_preds.append(weighted_pred)\n            all_confidences.append(avg_conf)\n    \n    final_predictions = torch.cat(all_preds, dim=0)\n    final_confidences = torch.cat(all_confidences, dim=0)\n    return final_predictions, final_confidences\n\n# =============================================================================\n# 4. Training Components and Evaluation\n# =============================================================================\n\ndef tm_score_loss(pred, target, confidence=None):\n    target = target.permute(0, 2, 1, 3)\n    pred_struct = pred[:, :, 0, :]\n    all_tm_scores = []\n    for i in range(target.shape[2]):\n        target_struct = target[:, :, i, :]\n        squared_dists = torch.sum((pred_struct - target_struct)**2, dim=-1) + 1e-8\n        dists = torch.sqrt(squared_dists)\n        seq_len = target.shape[1]\n        d0 = 1.24 * (seq_len - 15)**(1/3) - 1.8\n        tm_score_components = 1 / (1 + (dists / d0)**2)\n        if confidence is not None:\n            norm_confidence = seq_len * confidence / confidence.sum(dim=1, keepdim=True)\n            tm_score_components = tm_score_components * norm_confidence\n        tm_scores = tm_score_components.mean(dim=1)\n        all_tm_scores.append(tm_scores)\n    all_tm_scores = torch.stack(all_tm_scores, dim=1)\n    best_tm_scores = all_tm_scores.max(dim=1)[0]\n    l2_reg = torch.mean(torch.norm(pred_struct, dim=2)) * 0.001\n    return -best_tm_scores.mean() + l2_reg\n\ndef compute_accuracy(pred, target, threshold=1.0):\n    pred_struct = pred[:, :, 0, :]\n    target = target.permute(0, 2, 1, 3)\n    best_acc = torch.zeros(pred.size(0), device=pred.device)\n    for i in range(target.shape[2]):\n        target_struct = target[:, :, i, :]\n        distances = torch.sqrt(torch.sum((pred_struct - target_struct)**2, dim=-1) + 1e-8)\n        correct = (distances < threshold).float().mean(dim=1)\n        best_acc = torch.max(torch.stack([best_acc, correct]), dim=0)[0]\n    return best_acc.mean().item()\n\n# =============================================================================\n# 5. Main Training Pipeline and Experiment Management\n# =============================================================================\n\ndef train_and_validate(model, train_loader, val_loader, epochs=EPOCHS, patience=10, save_best=True, model_path='best_model.pth'):\n    optimizer = optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-5)\n    scheduler = ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=5, verbose=True)\n    best_val_loss = float('inf')\n    best_val_acc = 0\n    no_improve_epochs = 0\n    history = {'epoch': [], 'train_loss': [], 'val_loss': [], 'val_acc': []}\n    \n    for epoch in range(epochs):\n        print(f\"\\nEpoch {epoch+1}/{epochs}\")\n        model.train()\n        total_loss = 0\n        total_samples = 0\n        progress_bar = tqdm(train_loader, desc=\"Training\")\n        train_accuracies = []\n        for features, targets in progress_bar:\n            features = augment_features(features, noise_std=0.01).to(DEVICE)\n            targets = targets.to(DEVICE)\n            optimizer.zero_grad()\n            structure_pred, confidence = model(features)\n            loss = tm_score_loss(structure_pred, targets, confidence)\n            combined_loss = loss\n            combined_loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            optimizer.step()\n            batch_size = features.size(0)\n            total_loss += combined_loss.item() * batch_size\n            total_samples += batch_size\n            train_acc = compute_accuracy(structure_pred, targets)\n            train_accuracies.append(train_acc)\n            progress_bar.set_postfix({'loss': combined_loss.item(), 'avg_loss': total_loss / total_samples})\n        \n        epoch_train_loss = total_loss / total_samples\n        epoch_train_acc = np.mean(train_accuracies)\n        \n        model.eval()\n        val_loss = 0\n        val_accs = []\n        with torch.no_grad():\n            for features, targets in val_loader:\n                features, targets = features.to(DEVICE), targets.to(DEVICE)\n                structure_pred, confidence = model(features)\n                loss = tm_score_loss(structure_pred, targets, confidence)\n                val_loss += loss.item() * features.size(0)\n                val_accs.append(compute_accuracy(structure_pred, targets))\n        epoch_val_loss = val_loss / len(val_loader.dataset)\n        epoch_val_acc = np.mean(val_accs)\n        \n        history['epoch'].append(epoch+1)\n        history['train_loss'].append(epoch_train_loss)\n        history['val_loss'].append(epoch_val_loss)\n        history['val_acc'].append(epoch_val_acc)\n        \n        scheduler.step(epoch_val_loss)\n        print(f\"Train Loss: {epoch_train_loss:.4f}, Val Loss: {epoch_val_loss:.4f}, Val Accuracy: {epoch_val_acc:.4f}\")\n        \n        if epoch_val_loss < best_val_loss:\n            best_val_loss = epoch_val_loss\n            best_val_acc = epoch_val_acc\n            no_improve_epochs = 0\n            print(f\"New best validation loss: {best_val_loss:.4f}\")\n            if save_best:\n                torch.save({\n                    'epoch': epoch,\n                    'model_state_dict': model.state_dict(),\n                    'optimizer_state_dict': optimizer.state_dict(),\n                    'val_loss': epoch_val_loss,\n                    'val_acc': epoch_val_acc,\n                }, model_path)\n                print(f\"Model saved to {model_path}\")\n        else:\n            no_improve_epochs += 1\n            print(f\"No improvement for {no_improve_epochs} epochs.\")\n            if no_improve_epochs >= patience:\n                print(f\"Early stopping after {epoch+1} epochs.\")\n                break\n    \n    # Plot training history\n    plt.figure(figsize=(10,6))\n    plt.plot(history['epoch'], history['train_loss'], label='Train Loss')\n    plt.plot(history['epoch'], history['val_loss'], label='Validation Loss')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.title('Training and Validation Loss History')\n    plt.legend()\n    plt.grid(True)\n    plt.savefig(\"training_loss_history.png\")\n    plt.show()\n    \n    plt.figure(figsize=(10,6))\n    plt.plot(history['epoch'], history['val_acc'], label='Validation Accuracy', marker='o')\n    plt.xlabel('Epoch')\n    plt.ylabel('Accuracy')\n    plt.title('Validation Accuracy History')\n    plt.legend()\n    plt.grid(True)\n    plt.savefig(\"validation_accuracy_history.png\")\n    plt.show()\n    \n    return best_val_loss, best_val_acc, history\n\ndef test_prediction(model, test_loader, sample_submission, model_path='best_model.pth'):\n    checkpoint = torch.load(model_path)\n    model.load_state_dict(checkpoint['model_state_dict'])\n    print(f\"Loaded model from epoch {checkpoint['epoch']} with val_loss {checkpoint['val_loss']:.4f}\")\n    model.eval()\n    predictions = []\n    target_ids = []\n    \n    for batch_idx, (features, _) in enumerate(tqdm(test_loader, desc=\"Generating predictions\")):\n        features = features.to(DEVICE)\n        with torch.no_grad():\n            structure_pred, _ = model(features)\n        pred_coords = structure_pred[:, :, 0, :].cpu().numpy()\n        for i in range(len(pred_coords)):\n            idx = batch_idx * test_loader.batch_size + i\n            target_id = test_loader.dataset.seq_df.iloc[idx]['target_id']\n            target_ids.append(target_id)\n            predictions.append(pred_coords[i])\n    \n    submission_rows = []\n    for target_id, pred_coords in zip(target_ids, predictions):\n        row = {'ID': target_id}\n        for i, (x, y, z) in enumerate(pred_coords, 1):\n            if i > 100:\n                break\n            row[f'x_{i}'] = x\n            row[f'y_{i}'] = y\n            row[f'z_{i}'] = z\n        submission_rows.append(row)\n    \n    submission_df = pd.DataFrame(submission_rows)\n    return submission_df\n\ndef analyze_model_performance(model, val_loader, output_dir=\"./analysis\"):\n    os.makedirs(output_dir, exist_ok=True)\n    model.eval()\n    features, targets = next(iter(val_loader))\n    features, targets = features.to(DEVICE), targets.to(DEVICE)\n    with torch.no_grad():\n        structure_pred, confidence = model(features)\n    plt.figure(figsize=(10, 6))\n    plt.hist(confidence.cpu().numpy().flatten(), bins=50)\n    plt.title(\"Confidence Distribution\")\n    plt.xlabel(\"Confidence\")\n    plt.ylabel(\"Count\")\n    plt.savefig(os.path.join(output_dir, \"confidence_distribution.png\"))\n    plt.show()\n    plt.close()\n    sample_idx = 0\n    pred_struct = structure_pred[sample_idx, :, 0, :].cpu().numpy()\n    targets_permuted = targets.permute(0, 2, 1, 3)\n    true_struct = targets_permuted[sample_idx, :, 0, :].cpu().numpy()\n    fig = plt.figure(figsize=(12, 10))\n    ax = fig.add_subplot(111, projection='3d')\n    ax.scatter(pred_struct[:, 0], pred_struct[:, 1], pred_struct[:, 2],\n               c='blue', marker='o', alpha=0.6, label='Predicted')\n    ax.scatter(true_struct[:, 0], true_struct[:, 1], true_struct[:, 2],\n               c='red', marker='^', alpha=0.6, label='True')\n    for i in range(len(pred_struct)-1):\n        ax.plot([pred_struct[i, 0], pred_struct[i+1, 0]],\n                [pred_struct[i, 1], pred_struct[i+1, 1]],\n                [pred_struct[i, 2], pred_struct[i+1, 2]], 'b-', alpha=0.3)\n        ax.plot([true_struct[i, 0], true_struct[i+1, 0]],\n                [true_struct[i, 1], true_struct[i+1, 1]],\n                [true_struct[i, 2], true_struct[i+1, 2]], 'r-', alpha=0.3)\n    ax.set_title(\"3D Structure Comparison\")\n    ax.legend()\n    plt.savefig(os.path.join(output_dir, \"structure_comparison_3d.png\"))\n    plt.show()\n    plt.close()\n    try:\n        background = features[:10].cpu()\n        explainer = shap.DeepExplainer(model, background)\n        shap_values = explainer.shap_values(features[:5].cpu())\n        plt.figure(figsize=(12, 8))\n        shap.summary_plot(shap_values, features[:5].cpu(), feature_names=['A', 'C', 'G', 'U', 'Conservation', 'MSA Depth'])\n        plt.savefig(os.path.join(output_dir, \"shap_feature_importance.png\"))\n        plt.show()\n        plt.close()\n    except Exception as e:\n        print(f\"SHAP analysis failed: {e}\")\n\ndef evaluate_holdout(model, seq_df, label_df, msa_dir, holdout_fraction=0.2):\n    \"\"\"Evaluate model on an unseen holdout set from the training data.\"\"\"\n    train_df, holdout_df = train_test_split(seq_df, test_size=holdout_fraction, random_state=42)\n    holdout_dataset = RNADataset(holdout_df, label_df, msa_dir, augment=False)\n    holdout_loader = DataLoader(holdout_dataset, batch_size=BATCH_SIZE, shuffle=False, collate_fn=custom_collate_fn, num_workers=0)\n    model.eval()\n    holdout_losses = []\n    holdout_accuracies = []\n    with torch.no_grad():\n        for features, targets in holdout_loader:\n            features, targets = features.to(DEVICE), targets.to(DEVICE)\n            structure_pred, confidence = model(features)\n            loss = tm_score_loss(structure_pred, targets, confidence)\n            holdout_losses.append(loss.item())\n            holdout_accuracies.append(compute_accuracy(structure_pred, targets))\n    avg_loss = np.mean(holdout_losses)\n    avg_acc = np.mean(holdout_accuracies)\n    print(f\"Holdout Evaluation - Loss: {avg_loss:.4f}, Accuracy: {avg_acc:.4f}\")\n    return avg_loss, avg_acc\n\ndef cross_validation(seq_df, label_df, msa_dir, n_folds=5, model_class=ImprovedRNA3DModel, **model_kwargs):\n    kf = StratifiedKFold(n_splits=n_folds, shuffle=True, random_state=42)\n    val_losses = []\n    val_accs = []\n    strata = seq_df['sequence'].apply(len)\n    strata = pd.qcut(strata, 5, labels=False)\n    for fold, (train_idx, val_idx) in enumerate(kf.split(seq_df, strata)):\n        print(f\"\\n--- Fold {fold+1}/{n_folds} ---\")\n        train_seq_df = seq_df.iloc[train_idx].reset_index(drop=True)\n        val_seq_df = seq_df.iloc[val_idx].reset_index(drop=True)\n        train_dataset = RNADataset(train_seq_df, label_df, msa_dir, augment=True)\n        val_dataset = RNADataset(val_seq_df, label_df, msa_dir, augment=False)\n        train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, collate_fn=custom_collate_fn, num_workers=0)\n        val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, collate_fn=custom_collate_fn, num_workers=0)\n        model = model_class(**model_kwargs).to(DEVICE)\n        fold_val_loss, fold_val_acc = train_and_validate(model, train_loader, val_loader, model_path=f'model_fold{fold+1}.pth')\n        val_losses.append(fold_val_loss)\n        val_accs.append(fold_val_acc)\n        print(f\"Fold {fold+1} validation loss: {fold_val_loss:.4f}, accuracy: {fold_val_acc:.4f}\")\n        del model, train_dataset, val_dataset, train_loader, val_loader\n        clear_memory()\n    print(\"\\n--- Cross-Validation Results ---\")\n    print(f\"Mean validation loss: {np.mean(val_losses):.4f} ± {np.std(val_losses):.4f}\")\n    print(f\"Mean validation accuracy: {np.mean(val_accs):.4f} ± {np.std(val_accs):.4f}\")\n    return val_losses, val_accs\n\n# =============================================================================\n# 6. Main Execution Pipeline\n# =============================================================================\n\ndef main():\n    print(f\"Using device: {DEVICE}\")\n    train_seqs, val_seqs, test_seqs, train_labels, val_labels, sample_submission = load_data()\n    verify_data_splits(train_seqs, val_seqs, test_seqs)\n    print(f\"Training on {len(train_seqs)} sequences, validating on {len(val_seqs)}, testing on {len(test_seqs)}\")\n    msa_dir = os.path.join(DATA_PATH, \"MSA\")\n    \n    # Create datasets\n    train_dataset = RNADataset(train_seqs, train_labels, msa_dir, augment=True)\n    val_dataset = RNADataset(val_seqs, val_labels, msa_dir, augment=False)\n    test_dataset = RNADataset(test_seqs, pd.DataFrame(), msa_dir, augment=False)\n    \n    # Create data loaders with custom collate function (fixed sequence length)\n    train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, collate_fn=custom_collate_fn, num_workers=0)\n    val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, collate_fn=custom_collate_fn, num_workers=0)\n    test_loader = DataLoader(test_dataset, batch_size=BATCH_SIZE, shuffle=False, collate_fn=custom_collate_fn, num_workers=0)\n    \n    # Initialize model (improved architecture)\n    model = ImprovedRNA3DModel(input_channels=6, seq_length=SEQ_LENGTH, num_structures=NUM_PREDICTIONS).to(DEVICE)\n    \n    # Train model and record history\n    best_val_loss, best_val_acc, history = train_and_validate(\n        model, train_loader, val_loader, epochs=EPOCHS, patience=10, save_best=True, model_path='best_rna_model.pth'\n    )\n    print(f\"Training complete. Best validation loss: {best_val_loss:.4f}, accuracy: {best_val_acc:.4f}\")\n    \n    # Evaluate on unseen holdout data\n    print(\"Evaluating on holdout data...\")\n    holdout_loss, holdout_acc = evaluate_holdout(model, train_seqs, train_labels, msa_dir, holdout_fraction=0.2)\n    \n    # Graphical plots for overall peaked accuracies (train, validation, holdout)\n    plt.figure(figsize=(10,6))\n    plt.plot(history['epoch'], history['train_loss'], label=\"Train Loss\", marker='o')\n    plt.plot(history['epoch'], history['val_loss'], label=\"Validation Loss\", marker='o')\n    plt.axhline(y=holdout_loss, color='red', linestyle='--', label=f\"Holdout Loss: {holdout_loss:.4f}\")\n    plt.xlabel(\"Epoch\")\n    plt.ylabel(\"Loss\")\n    plt.title(\"Training, Validation, and Holdout Loss\")\n    plt.legend()\n    plt.grid(True)\n    plt.savefig(\"overall_loss_history.png\")\n    plt.show()\n    \n    plt.figure(figsize=(10,6))\n    plt.plot(history['epoch'], history['val_acc'], label=\"Validation Accuracy\", marker='o')\n    plt.axhline(y=holdout_acc, color='green', linestyle='--', label=f\"Holdout Accuracy: {holdout_acc:.4f}\")\n    plt.xlabel(\"Epoch\")\n    plt.ylabel(\"Accuracy\")\n    plt.title(\"Validation and Holdout Accuracy\")\n    plt.legend()\n    plt.grid(True)\n    plt.savefig(\"overall_accuracy_history.png\")\n    plt.show()\n    \n    # Analyze model performance (feature importance, SHAP, etc.)\n    analyze_model_performance(model, val_loader)\n    \n    # Generate predictions with Test-Time Augmentation\n    print(\"Generating predictions with TTA...\")\n    tta_predictions, tta_confidences = tta_predict(model, test_loader, tta_iterations=5)\n    \n    # Create submission file from TTA predictions\n    submission_df = test_prediction(model, test_loader, sample_submission, model_path='best_rna_model.pth')\n    submission_df.to_csv(\"submission.csv\", index=False)\n    print(\"Submission file created: submission.csv\")\n    \n    # Optional cross-validation (if desired)\n    if False:\n        print(\"\\nRunning cross-validation...\")\n        cv_losses, cv_accs = cross_validation(\n            train_seqs, train_labels, msa_dir, n_folds=5, model_class=ImprovedRNA3DModel,\n            input_channels=6, seq_length=SEQ_LENGTH, num_structures=NUM_PREDICTIONS\n        )\n    \n    print(\"=\"*50)\n    print(\"RNA 3D Structure Prediction Pipeline Complete\")\n    print(\"=\"*50)\n    clear_memory()\n    print(\"All resources cleaned up. Execution complete.\")\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T12:58:54.887963Z","iopub.execute_input":"2025-04-16T12:58:54.888526Z","iopub.status.idle":"2025-04-16T13:20:41.724467Z","shell.execute_reply.started":"2025-04-16T12:58:54.888494Z","shell.execute_reply":"2025-04-16T13:20:41.723643Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Summary\n*******************************************************************************************\n# 7 . Explanations and Interpretations of the Attached Plots\n\nThe attached image includes several figures generated by the code. Each subplot showcases a different aspect of training, validation, or inference outcomes for the RNA 3D structure prediction model. Below is a detailed explanation of each figure, what it represents, and how to interpret it.\nA. Training and Validation Loss History\n\n    What It Shows:\n    This figure plots two lines: the training loss (in orange) and the validation loss (in blue) over a series of epochs.\n\n    Interpretation:\n\n        Training Loss Curve: Typically starts high (since the model parameters are initially random or untrained) and decreases over epochs as the model learns patterns in the training data.\n\n        Validation Loss Curve: Monitors how well the model generalizes to unseen data. Ideally, it should follow the training loss downward, but might plateau or increase if overfitting occurs.\n\n        Observations:\n\n            If the training loss continues to decrease but the validation loss starts to rise, that indicates overfitting.\n\n            If both training and validation losses converge, that suggests the model is learning a robust representation.\n\n# 8 . Validation Accuracy History\n\n    What It Shows:\n    A single curve (blue line with dots) depicting the validation accuracy by epoch.\n\n    Interpretation:\n\n        Accuracy Metric: Accuracy in this specific pipeline is computed by checking how many coordinates are predicted within a certain threshold distance (default = 1.0 Å) from the true coordinates, then averaging the best alignment across multiple structure predictions.\n\n        Progress Over Epochs:\n\n            Early in training, accuracy may be low because the model’s predictions are not yet close to the actual 3D coordinates.\n\n            As training progresses, we expect accuracy to improve if the model successfully learns the mapping from sequence to structure.\n\n        Observations: A sharp rise indicates that the model learns relatively quickly after enough epochs. A plateau means the model is no longer improving.\n\n# 9. Training, Validation, and Holdout Loss\n\n    What It Shows:\n    This plot typically displays three elements:\n\n        Training Loss vs. Epoch\n\n        Validation Loss vs. Epoch\n\n        A dashed horizontal line (in red) indicating the holdout loss\n\n    Interpretation:\n\n        Holdout Set: This is a portion of the data not used in the main training or validation sets. It’s intended to simulate “real-world” unseen data, giving another check on model generalization.\n\n        Comparison:\n\n            If the holdout loss is significantly higher than the validation loss, it might mean the validation set is easier or not fully representative of real unseen data.\n\n            A close alignment between the holdout loss and validation loss suggests stable generalization.\n\n# 10. Validation and Holdout Accuracy\n\n    What It Shows:\n    A line (blue) for validation accuracy across epochs, plus a dashed horizontal line (green) showing holdout accuracy.\n\n    Interpretation:\n\n        Validation vs. Holdout: The holdout accuracy is a single value measured once (or occasionally) after training.\n\n        Assessing Stability:\n\n            If the holdout accuracy is similar to the maximum validation accuracy, it indicates that the model’s performance is consistent.\n\n            A large gap (e.g., high validation accuracy but low holdout accuracy) could mean the model is overfitting to the validation set.\n\n# 11. . Confidence Distribution\n\n    What It Shows:\n    A histogram of the confidence scores output by the confidence head in the improved model.\n\n    Interpretation:\n\n        Confidence Scores: Each residue in the predicted structure is assigned a value between 0 and 1. Higher values generally imply higher trust in that part of the prediction.\n\n        Distribution Shape:\n\n            A peak near 1.0 could indicate the model is very confident for most positions.\n\n            A broader distribution (ranging across 0 to 1) indicates varying confidence across positions.\n\n        Practical Usage: In test-time augmentation (TTA), these confidence scores can help weigh predictions from multiple forward passes, emphasizing more confident predictions.\n\n# 12. 3D Structure Comparison\n\n    What It Shows:\n    A 3D scatter plot of Predicted (blue circles) vs. True (red triangles) coordinates for a single sample. Lines connect consecutive residues, providing a “backbone trace.”\n\n    Interpretation:\n\n        Visual Spatial Alignment: Quickly see if the predicted structure (blue) roughly overlaps or aligns with the true structure (red).\n\n        Backbone Continuity: Consecutive residues are connected with lines (blue lines for the predicted structure, red lines for the true structure) to show the chain trace.\n\n        Looking for Deviation: Large spatial deviations between predicted and true points highlight where the model might need improvement. Good alignment indicates strong predictive capability.\n\n******************************************************************************************\nIn summary, the code implements a full-fledged pipeline for predicting the 3D structure of RNA molecules from sequence data. It meticulously covers:\n\n    Data preparation: Reading data, augmenting RNA sequences, and feature normalization.\n\n    Custom dataset handling: Integrating sequence and conservation features, along with processing MSA files and ground truth coordinate labels.\n\n    Model architecture: Starting from a simple convolutional network and evolving into a more complex and robust design that includes attention mechanisms, transformers, and graph convolution layers.\n\n    Training and evaluation: A training loop with early stopping, learning rate scheduling, loss function based on the TM-score, and evaluation using accuracy and visualization metrics.\n\n    Experiment management: Functions for holdout evaluation, cross-validation, test-time augmentation, and final prediction generation with submission file creation.\n\n    Resource management: Functions to optimize memory usage and ensure reproducibility.\n\nThis modular design makes the pipeline extensible and well-suited for experimentation and improvement in RNA 3D structure prediction tasks.","metadata":{}},{"cell_type":"markdown","source":"****************************************************************************************\nIf you found this code helpful, please consider upvoting and supporting the Handsonlabs Software Academy initiative. We are currently raising funds to establish a 1,000-capacity, well-equipped Computer-Based Testing and Skill Acquisition Center in Lagos, Nigeria. \n\nThis center will provide training for Nigerian youths to effectively use the internet and various computational resources, productively. For more information on how to support or to get in touch, please directly message us. Handsonlabs Software Academy is ID verified on the Kaggle platform. Handsonlabs can also be reached out to on YouTube, Facebook and WhatsApp respectively.\n******************************************************************************************\nAFRIMEX Mobile APP ID:  hands641 (name: Oluwatobi Owoeye) : \nFor Handsonlabs Software Academy\n*******************************************************************************************\nNOTE: This code is meant to help you all in this competition and has not been graded by the Kaggle Letterboard.","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}