{"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":12276181,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":397358,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":325826,"modelId":346679}],"dockerImageVersionId":31040,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Data Loader","metadata":{}},{"cell_type":"code","source":"# RNADataLoader Class\nimport pandas as pd\nimport numpy as np\nfrom sklearn.preprocessing import LabelEncoder\n\nclass RNADataLoader:\n    def __init__(self):\n        self.nucleotide_encoder = LabelEncoder()\n        self.nucleotide_encoder.fit(['A', 'U', 'G', 'C'])\n        \n    def one_hot_encode(self, sequence):\n        \"\"\"Convert RNA sequence to one-hot encoding\"\"\"\n        encoded = np.zeros((len(sequence), 4))\n        for i, nt in enumerate(sequence):\n            if nt in self.nucleotide_encoder.classes_:\n                encoded[i, self.nucleotide_encoder.transform([nt])[0]] = 1\n            else:\n                # Handle unexpected characters (e.g., '-', 'N') by skipping or encoding as zeros\n                print(f\"Warning: Unexpected nucleotide '{nt}' encountered. Encoding as zeros.\")\n        return encoded\n        \n    def load_train_data(self, seq_file, labels_file):\n        \"\"\"Load training data from CSV files\"\"\"\n        sequences_df = pd.read_csv(seq_file)\n        labels_df = pd.read_csv(labels_file)\n        \n        X = []\n        y = []\n        \n        for _, row in sequences_df.iterrows():\n            target_id = row['target_id']\n            sequence = row['sequence']\n            \n            # Get coordinates for this sequence\n            seq_labels = labels_df[labels_df['ID'].str.startswith(target_id)]  # Use \"ID\" for labels\n            \n            if len(seq_labels) > 0:\n                # Extract coordinates\n                coords = seq_labels[['x_1', 'y_1', 'z_1']].values\n                # Mask invalid values\n                coords[np.abs(coords) > 1e6] = 0.0  # or np.nan, but 0.0 is safe for masking\n                # Mask NaNs\n                coords[np.isnan(coords)] = 0.0\n                # One-hot encode the sequence\n                encoded_seq = self.one_hot_encode(sequence)\n                \n                X.append(encoded_seq)\n                y.append(coords)\n        \n        # Convert to arrays with padding to ensure all sequences have same length\n        max_length = max(len(x) for x in X)\n        X_padded = np.zeros((len(X), max_length, 4))\n        y_padded = np.zeros((len(y), max_length, 3))  # Ensure same max_length\n\n        for i, (x, coords) in enumerate(zip(X, y)):\n            X_padded[i, :len(x)] = x\n            y_padded[i, :len(coords)] = coords\n            \n        print(f\"X shape: {X_padded.shape}, y shape: {y_padded.shape}\")\n        return X_padded, y_padded, sequences_df['target_id'].values\n    \n    def load_test_data(self, seq_file):\n        \"\"\"Load test data from CSV file\"\"\"\n        sequences_df = pd.read_csv(seq_file)\n        \n        X = []\n        sequence_lengths = []\n        \n        for _, row in sequences_df.iterrows():\n            sequence = row['sequence']\n            encoded_seq = self.one_hot_encode(sequence)\n            X.append(encoded_seq)\n            sequence_lengths.append(len(sequence))\n        \n        # Pad sequences to max length\n        max_length = max(len(x) for x in X)\n        X_padded = np.zeros((len(X), max_length, 4))\n        \n        for i, x in enumerate(X):\n            X_padded[i, :len(x)] = x\n            \n        return X_padded, sequence_lengths, sequences_df['target_id'].values","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T17:11:52.724082Z","iopub.execute_input":"2025-05-16T17:11:52.724400Z","iopub.status.idle":"2025-05-16T17:11:52.736235Z","shell.execute_reply.started":"2025-05-16T17:11:52.724381Z","shell.execute_reply":"2025-05-16T17:11:52.735489Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model ","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\ndef create_base_pairing_mask(x):\n    \"\"\"Create a mask for valid RNA base pairs following established folding rules\n    References:\n    - Mathews et al. (2004) \"Incorporating chemical modification constraints into a dynamic programming algorithm for prediction of RNA secondary structure\"\n    - Zuker & Sankoff (1984) \"RNA secondary structures and their prediction\"\n    \"\"\"\n    batch_size, seq_len, _ = x.shape\n    \n    # Extract base type positions: [A,C,G,U]\n    A = x[:, :, 0].unsqueeze(2)\n    C = x[:, :, 1].unsqueeze(2)\n    G = x[:, :, 2].unsqueeze(2)\n    U = x[:, :, 3].unsqueeze(2)\n    \n    # Create pairing masks with different strengths\n    # Strong canonical pairs\n    A_U = torch.matmul(A, U.transpose(1,2)) * 0.9  # A-U pairs (strong)\n    U_A = torch.matmul(U, A.transpose(1,2)) * 0.9  # U-A pairs (strong)\n    G_C = torch.matmul(G, C.transpose(1,2))  # G-C pairs (strongest)\n    C_G = torch.matmul(C, G.transpose(1,2))  # C-G pairs (strongest)\n    \n    # Weak wobble pairs (G-U)\n    G_U = torch.matmul(G, U.transpose(1,2)) * 0.7  # G-U pairs (weaker)\n    U_G = torch.matmul(U, G.transpose(1,2)) * 0.7  # U-G pairs (weaker)\n    \n    # Combine all valid pairs with their respective strengths\n    pairing_mask = G_C + C_G + A_U + U_A + G_U + U_G\n    \n    # Create smooth distance penalty based on loop size constraints\n    positions = torch.arange(seq_len, device=x.device)\n    distances = torch.abs(positions.unsqueeze(1) - positions.unsqueeze(0))\n    \n    # Smooth penalty function:\n    # - Distance < 3: prohibited (physical constraint)\n    # - Distance 3-4: suboptimal but allowed\n    # - Distance 4-7: optimal hairpin size\n    # - Distance > 7: allowed but slightly penalized for long-range interactions\n    distance_penalty = torch.zeros_like(distances, dtype=torch.float)\n    distance_penalty = torch.where(distances < 3, torch.zeros_like(distances, dtype=torch.float), distance_penalty)\n    distance_penalty = torch.where((distances >= 3) & (distances < 4), 0.5 * torch.ones_like(distances, dtype=torch.float), distance_penalty)\n    distance_penalty = torch.where((distances >= 4) & (distances <= 7), torch.ones_like(distances, dtype=torch.float), distance_penalty)\n    distance_penalty = torch.where(distances > 7, 0.8 * torch.ones_like(distances, dtype=torch.float), distance_penalty)\n    \n    # Apply distance penalty and add small baseline for non-pairs\n    pairing_mask = pairing_mask * distance_penalty + 0.1\n    \n    return pairing_mask\n\ndef compute_distance_violations(coords, pairing_mask, min_distance=3.0, \n                             wc_pair_distance=5.9, backbone_distance=6.0):\n    \"\"\"\n    Compute distance violation penalties based on RNA structural constraints\n    References:\n    - Watson-Crick pair distance ~5.9Å (Regions et al. 2011)\n    - P-P backbone distance ~6.0Å (Richardson et al. 2008)\n    - Minimum allowed distance 3.0Å (van der Waals + water shell)\n    \"\"\"\n    batch_size, seq_length, _ = coords.size()\n    \n    # Compute all pairwise distances for each sequence in batch\n    # Shape: (batch_size, seq_length, seq_length)\n    distances = torch.cdist(coords, coords, p=2)\n    \n    # 1. Minimum distance violation (steric clash)\n    min_dist_violation = torch.relu(min_distance - distances)\n    identity_mask = torch.eye(seq_length, device=coords.device).unsqueeze(0)\n    min_dist_violation = min_dist_violation * (1 - identity_mask)  # Exclude self-distances\n    \n    # 2. Base pair distance violation (for paired bases)\n    # Ensure pairing_mask is expanded if needed\n    if len(pairing_mask.shape) == 2:\n        pairing_mask = pairing_mask.unsqueeze(0).expand(batch_size, -1, -1)\n    pair_dist_violation = torch.abs(distances - wc_pair_distance) * (pairing_mask > 0.5)\n    \n    # 3. Sequential backbone distance violation\n    # Create a mask for consecutive residues using diagonal shift\n    backbone_mask = torch.zeros(seq_length, seq_length, device=coords.device)\n    idx = torch.arange(seq_length-1, device=coords.device)\n    backbone_mask[idx, idx+1] = 1.0  # Set 1s on first diagonal above main diagonal\n    backbone_mask = backbone_mask.unsqueeze(0).expand(batch_size, -1, -1)\n    backbone_violation = torch.abs(distances - backbone_distance) * backbone_mask\n    \n    # Combine violations with weights\n    total_violation = (min_dist_violation * 2.0 +  # Stronger penalty for steric clashes\n                      pair_dist_violation * 1.0 +  # Base pair distance penalty\n                      backbone_violation * 1.0)    # Backbone distance penalty\n    \n    return total_violation.mean()  # Average over all violations\n\nclass RNAStructureModel(nn.Module):\n    def __init__(self, seq_length, feature_dim=4):\n        super(RNAStructureModel, self).__init__()\n        self.seq_length = seq_length\n        self.feature_dim = feature_dim\n        \n        # Feature extraction blocks with residual connections\n        self.features = nn.ModuleList([\n            nn.Sequential(\n                nn.Conv1d(in_channels=feature_dim, out_channels=64, kernel_size=3, padding=1),\n                nn.BatchNorm1d(64),\n                nn.ReLU()\n            ),\n            nn.Sequential(\n                nn.Conv1d(in_channels=64, out_channels=64, kernel_size=3, padding=1),\n                nn.BatchNorm1d(64)\n            ),\n            nn.Sequential(\n                nn.Conv1d(in_channels=64, out_channels=128, kernel_size=3, padding=1),\n                nn.BatchNorm1d(128),\n                nn.ReLU()\n            ),\n            nn.Sequential(\n                nn.Conv1d(in_channels=128, out_channels=128, kernel_size=3, padding=1),\n                nn.BatchNorm1d(128)\n            )\n        ])\n        \n        # Projection layers for residual connections\n        self.project1 = nn.Conv1d(in_channels=feature_dim, out_channels=64, kernel_size=1)\n        self.project2 = nn.Conv1d(in_channels=64, out_channels=128, kernel_size=1)\n        \n        # Batch normalization after residual connections\n        self.bn_after_res1 = nn.BatchNorm1d(64)\n        self.bn_after_res2 = nn.BatchNorm1d(128)\n        \n        # LSTM layers organized using LSTMBlock\n        self.lstm1 = LSTMBlock(input_size=128, hidden_size=128, bidirectional=True, layer_norm=True)\n        self.lstm2 = LSTMBlock(input_size=256, hidden_size=128, bidirectional=True, layer_norm=True)\n        \n        # Self-attention layer\n        self.base_pair_attention = BaseAwareAttention(embed_dim=256, num_heads=4)\n        \n        # Output layer for 3D coordinates (x, y, z) for each position\n        self.output_layer = nn.Linear(256, 3)\n\n    def forward(self, x, return_distance_loss=False):\n        \"\"\"Forward pass of the model\"\"\"\n        pairing_mask = create_base_pairing_mask(x)\n\n        batch_size, seq_len, feat_dim = x.shape\n        \n        # Transpose for Conv1d: [batch, seq_len, features] -> [batch, features, seq_len]\n        x = x.transpose(1, 2)\n        \n        # Save input for residual connection\n        identity1 = self.project1(x)\n        \n        # First conv block\n        x = self.features[0](x)\n        x = self.features[1](x)\n        x = x + identity1\n        x = self.bn_after_res1(x)\n        x = F.relu(x)\n        \n        # Second conv block\n        identity2 = self.project2(x)\n        x = self.features[2](x)\n        x = self.features[3](x)\n        x = x + identity2\n        x = self.bn_after_res2(x)\n        x = F.relu(x)\n        \n        # Transpose back for sequence processing\n        x = x.transpose(1, 2)\n        \n        # LSTM layers\n        x = self.lstm1(x)\n        x = self.lstm2(x)\n        \n        # Self-attention with memory optimization\n        x = self.base_pair_attention(x, pairing_mask)\n        \n        # Final output layer\n        coords = self.output_layer(x)\n        \n        if return_distance_loss:\n            distance_loss = compute_distance_violations(coords, pairing_mask)\n            return coords, distance_loss\n            \n        return coords\n    \n    def predict(self, x, sequence_lengths):\n        self.eval()  # Ensure the model is in evaluation mode\n        with torch.no_grad():\n            # Forward pass\n            output = self.forward(x)\n            \n            # If the forward pass returns a tuple, extract the predictions\n            if isinstance(output, tuple):\n                predictions = output[0]  # Extract the first element (coords)\n            else:\n                predictions = output  # If it's already a tensor\n\n            # Trim predictions based on sequence lengths\n            trimmed_predictions = []\n            for i, length in enumerate(sequence_lengths):\n                trimmed_predictions.append(predictions[i, :length].cpu().numpy())\n            \n            return trimmed_predictions\n    \n    def save(self, filepath):\n        \"\"\"Save the model weights\"\"\"\n        torch.save(self.state_dict(), filepath)\n    \n    def load(self, filepath):\n        \"\"\"Load model weights\"\"\"\n        self.load_state_dict(torch.load(filepath))\n\nclass LSTMBlock(nn.Module):\n    def __init__(self, input_size, hidden_size, bidirectional=True, layer_norm=False):\n        super(LSTMBlock, self).__init__()\n        self.lstm = nn.LSTM(input_size=input_size, hidden_size=hidden_size, \n                            batch_first=True, bidirectional=bidirectional)\n        self.layer_norm = nn.LayerNorm(hidden_size * 2) if bidirectional and layer_norm else None\n\n    def forward(self, x):\n        x, _ = self.lstm(x)  # LSTM returns (output, (h_n, c_n))\n        if self.layer_norm:\n            x = self.layer_norm(x)  # Apply layer normalization if enabled\n        return x\n\nclass BaseAwareAttention(nn.Module):\n    def __init__(self, embed_dim, num_heads):\n        super().__init__()\n        self.num_heads = num_heads\n        self.attention = nn.MultiheadAttention(embed_dim, num_heads, batch_first=True)\n        \n    def forward(self, x, base_pair_mask, padding_mask=None):\n        # Get the batch size and sequence length\n        batch_size, seq_length, _ = x.shape\n        \n        # Repeat the base_pair_mask for the number of attention heads\n        base_pair_mask = base_pair_mask.unsqueeze(1).repeat(1, self.num_heads, 1, 1)  # Shape: (batch_size, num_heads, seq_length, seq_length)\n        base_pair_mask = base_pair_mask.view(-1, seq_length, seq_length)  # Flatten to match attention batch size\n        \n        # Ensure base_pair_mask is of type float (required for attn_mask)\n        base_pair_mask = base_pair_mask.to(dtype=torch.float32)\n        \n        # Ensure padding_mask is of type float and repeat it for the number of attention heads\n        if padding_mask is not None:\n            padding_mask = padding_mask.to(dtype=torch.float32)\n            padding_mask = padding_mask.unsqueeze(1).repeat(1, self.num_heads, 1)  # Shape: (batch_size, num_heads, seq_length)\n            padding_mask = padding_mask.view(-1, seq_length)  # Flatten to match the batch size of base_pair_mask\n        \n        # Combine base_pair_mask and padding_mask\n        if padding_mask is not None:\n            combined_mask = base_pair_mask + padding_mask.unsqueeze(1).to(x.device)  # Combine masks\n        else:\n            combined_mask = base_pair_mask\n        \n        # Apply the combined mask during attention\n        attn_output, _ = self.attention(\n            x, x, x,\n            key_padding_mask=None,  # Set to None since we're using attn_mask\n            attn_mask=combined_mask.to(x.device)  # Use combined mask for attn_mask\n        )\n        return attn_output\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T17:11:52.737236Z","iopub.execute_input":"2025-05-16T17:11:52.737537Z","iopub.status.idle":"2025-05-16T17:11:52.767457Z","shell.execute_reply.started":"2025-05-16T17:11:52.737511Z","shell.execute_reply":"2025-05-16T17:11:52.766164Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Prediction","metadata":{}},{"cell_type":"code","source":"# Prediction and evaluation utilities\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\n\n# --- Prediction and evaluation functions from predict.py ---\ndef predict_structures(model_path, test_seq_file, output_path):\n    \"\"\"Predict RNA 3D structures for test sequences\"\"\"\n    data_loader = RNADataLoader()\n    print(\"Loading test data...\")\n    X, sequence_lengths, ids = data_loader.load_test_data(test_seq_file)\n    print(f\"Loaded {len(X)} sequences for prediction\")\n    X_tensor = torch.tensor(X, dtype=torch.float32)\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    X_tensor = X_tensor.to(device)\n    print(\"Loading model...\")\n    model = RNAStructureModel(seq_length=X.shape[1], feature_dim=X.shape[2])\n    checkpoint = torch.load(model_path, map_location=device)\n    if \"model_state_dict\" in checkpoint:\n        print(\"Loading model state_dict\")\n        model.load_state_dict(checkpoint[\"model_state_dict\"])\n    else:\n        model.load_state_dict(checkpoint)\n    model = model.to(device)\n    model.eval()\n    print(\"Predicting structures...\")\n    with torch.no_grad():\n        predictions = model.predict(X_tensor, sequence_lengths)\n\n    # Add 4 noisy versions for each prediction\n    all_rows = []\n    for seq_id, seq_len, pred in zip(ids, sequence_lengths, predictions):\n        if not hasattr(predict_structures, '_seq_cache'):\n            seq_df = pd.read_csv(test_seq_file)\n            predict_structures._seq_cache = dict(zip(seq_df['target_id'], seq_df['sequence']))\n        sequence = predict_structures._seq_cache.get(seq_id, '')\n        for pos in range(seq_len):\n            resname = sequence[pos] if pos < len(sequence) else ''\n            resid = pos + 1\n            # ID should be target_id_position (e.g., R1107_1)\n            id_with_pos = f\"{seq_id}_{resid}\"\n            row = {\n                'ID': id_with_pos,\n                'resname': resname,\n                'resid': resid,\n                'x_1': pred[pos][0],\n                'y_1': pred[pos][1],\n                'z_1': pred[pos][2],\n            }\n            for i in range(4):\n                noisy = pred + np.random.normal(0, 0.1, pred.shape)\n                row[f'x_{i+2}'] = noisy[pos][0]\n                row[f'y_{i+2}'] = noisy[pos][1]\n                row[f'z_{i+2}'] = noisy[pos][2]\n            all_rows.append(row)\n    output_df = pd.DataFrame(all_rows)\n    cols = ['ID', 'resname', 'resid'] + [f'{c}_{i}' for i in range(1,6) for c in ['x','y','z']]\n    output_df = output_df[[c for c in cols if c in output_df.columns]]\n    os.makedirs(os.path.dirname(output_path), exist_ok=True) if os.path.dirname(output_path) else None\n    output_df.to_csv(output_path, index=False)\n    print(f\"Predictions saved to {output_path}\")\n    return output_df\n\ndef evaluate_predictions(predictions_file, ground_truth_file):\n    \"\"\"Evaluate predictions against ground truth\"\"\"\n    predictions = pd.read_csv(predictions_file)\n    ground_truth = pd.read_csv(ground_truth_file)\n    merged = predictions.merge(\n        ground_truth,\n        left_on=['target_id', 'position'],\n        right_on=['ID', 'position'],\n        suffixes=('_pred', '')\n    )\n    merged['squared_diff_x'] = (merged['x_1_pred'] - merged['x_1'])**2\n    merged['squared_diff_y'] = (merged['y_1_pred'] - merged['y_1'])**2\n    merged['squared_diff_z'] = (merged['z_1_pred'] - merged['z_1'])**2\n    merged['rmsd'] = np.sqrt(\n        merged['squared_diff_x'] + \n        merged['squared_diff_y'] + \n        merged['squared_diff_z']\n    )\n    overall_rmsd = np.sqrt(merged[['squared_diff_x', 'squared_diff_y', 'squared_diff_z']].sum().sum() / len(merged))\n    print(f\"Overall RMSD: {overall_rmsd:.4f}Å\")\n    sequence_rmsd = merged.groupby('target_id').apply(\n        lambda x: np.sqrt(\n            (x['squared_diff_x'].sum() + \n             x['squared_diff_y'].sum() + \n             x['squared_diff_z'].sum()) / len(x)\n        )\n    )\n    print(f\"Mean sequence RMSD: {sequence_rmsd.mean():.4f}Å\")\n    print(f\"Min sequence RMSD: {sequence_rmsd.min():.4f}Å\")\n    print(f\"Max sequence RMSD: {sequence_rmsd.max():.4f}Å\")\n    return overall_rmsd, sequence_rmsd\n\ndef get_sequence_labels(labels_df, target_id):\n    \"\"\"Get sequence labels for a specific target ID\"\"\"\n    seq_labels = labels_df[labels_df['ID'] == target_id]\n    return seq_labels\n\n# --- Run prediction and output submission.csv ---\n# Set your model and test file paths here\nmodel_path = \"/kaggle/input/phmfold650/pytorch/default/1/rna_structure_model_l5_lr005_epoch_650.pth\"  # Update as needed\ntest_seq_file = \"/kaggle/input/stanford-rna-3d-folding/test_sequences.csv\"  # Example: use train_sequences.csv as test input\noutput_path = \"/kaggle/working/submission.csv\"\n\n# Run prediction\ndf_pred = predict_structures(model_path, test_seq_file, output_path)\nprint(df_pred.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T17:11:52.768583Z","iopub.execute_input":"2025-05-16T17:11:52.768850Z","iopub.status.idle":"2025-05-16T17:11:55.023665Z","shell.execute_reply.started":"2025-05-16T17:11:52.768828Z","shell.execute_reply":"2025-05-16T17:11:55.022917Z"}},"outputs":[],"execution_count":null}]}