{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","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":11228175,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":219146636,"sourceType":"kernelVersion"}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q --no-index --find-links=/kaggle/input/pip-install-pyg torch_geometric","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:01:19.943825Z","iopub.execute_input":"2025-03-05T09:01:19.944071Z","iopub.status.idle":"2025-03-05T09:01:24.944485Z","shell.execute_reply.started":"2025-03-05T09:01:19.944038Z","shell.execute_reply":"2025-03-05T09:01:24.943423Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport torch\nfrom torch_geometric.data import Data\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch_geometric.nn import SAGEConv\nfrom torch_geometric.data import Data, DataLoader\nfrom sklearn.model_selection import train_test_split\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nimport random\n\n# Set random seeds for reproducibility\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    torch.backends.cudnn.deterministic = True\n\nset_seed(42)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:01:33.241840Z","iopub.execute_input":"2025-03-05T09:01:33.242174Z","iopub.status.idle":"2025-03-05T09:01:44.177902Z","shell.execute_reply.started":"2025-03-05T09:01:33.242130Z","shell.execute_reply":"2025-03-05T09:01:44.176930Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MAX_SEQ_LEN = 1024","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:01:44.178996Z","iopub.execute_input":"2025-03-05T09:01:44.179507Z","iopub.status.idle":"2025-03-05T09:01:44.183066Z","shell.execute_reply.started":"2025-03-05T09:01:44.179473Z","shell.execute_reply":"2025-03-05T09:01:44.182225Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define paths to data files\nTRAIN_SEQUENCES_PATH = \"/kaggle/input/stanford-rna-3d-folding/train_sequences.csv\"\nTRAIN_LABELS_PATH = \"/kaggle/input/stanford-rna-3d-folding/train_labels.csv\"\n# Load data\ntrain_sequences = pd.read_csv(TRAIN_SEQUENCES_PATH)\ntrain_labels = pd.read_csv(TRAIN_LABELS_PATH)\n\nprint(f\"Loaded {len(train_sequences)} RNA sequences and {len(train_labels)} nucleotide labels\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:01:44.184792Z","iopub.execute_input":"2025-03-05T09:01:44.185003Z","iopub.status.idle":"2025-03-05T09:01:44.505208Z","shell.execute_reply.started":"2025-03-05T09:01:44.184985Z","shell.execute_reply":"2025-03-05T09:01:44.504445Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Preprocess data\n# 1. Encoding nucleotides\nnucleotide_mapping = {'A': 0, 'C': 1, 'G': 2, 'U': 3}\nreverse_mapping = {0: 'A', 1: 'C', 2: 'G', 3: 'U'}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:01:44.506291Z","iopub.execute_input":"2025-03-05T09:01:44.506580Z","iopub.status.idle":"2025-03-05T09:01:44.510071Z","shell.execute_reply.started":"2025-03-05T09:01:44.506558Z","shell.execute_reply":"2025-03-05T09:01:44.509421Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 2. Create feature representation for each nucleotide\ndef one_hot_encode(nucleotide):\n    encoding = [0, 0, 0, 0]\n    if nucleotide in nucleotide_mapping:\n        encoding[nucleotide_mapping[nucleotide]] = 1\n    return encoding","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:01:44.510871Z","iopub.execute_input":"2025-03-05T09:01:44.511074Z","iopub.status.idle":"2025-03-05T09:01:44.530714Z","shell.execute_reply.started":"2025-03-05T09:01:44.511056Z","shell.execute_reply":"2025-03-05T09:01:44.530003Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Function to create a graph from an RNA sequence\ndef sequence_to_graph(sequence, target_id, labels_df=None, max_connections=MAX_SEQ_LEN):\n    \"\"\"\n    Create a graph representation of an RNA sequence.\n    \n    Args:\n        sequence: The RNA sequence\n        target_id: Identifier for the RNA\n        labels_df: Optional dataframe with 3D coordinate labels\n        max_connections: Maximum number of edges to create (to avoid CUDA OOM errors)\n        \n    Returns:\n        PyTorch Geometric Data object\n    \"\"\"\n    # One-hot encode each nucleotide\n    x = [one_hot_encode(nt) for nt in sequence]\n    x = torch.tensor(x, dtype=torch.float)\n    \n    # Create edges - connect adjacent nucleotides (backbone)\n    # and potentially other connections based on domain knowledge\n    edges = []\n    \n    # Always add backbone connections\n    for i in range(len(sequence) - 1):\n        # Connect to next nucleotide (backbone)\n        edges.append([i, i + 1])\n        edges.append([i + 1, i])  # Bidirectional\n    \n    # Add potential base-pairing connections, but limit total edges to avoid OOM\n    edge_count = len(edges)\n    max_additional_edges = max_connections - edge_count\n    \n    if max_additional_edges > 0:\n        potential_base_pairs = []\n        \n        # Identify potential base pairs (A-U, G-C)\n        for i in range(len(sequence)):\n            for j in range(i + 3, len(sequence)):  # Minimum loop size of 3\n                if (sequence[i] == 'A' and sequence[j] == 'U') or \\\n                   (sequence[i] == 'U' and sequence[j] == 'A') or \\\n                   (sequence[i] == 'G' and sequence[j] == 'C') or \\\n                   (sequence[i] == 'C' and sequence[j] == 'G'):\n                    # Store the potential base pair\n                    potential_base_pairs.append((i, j))\n        \n        # Randomly select base pairs if we have too many\n        if len(potential_base_pairs) > max_additional_edges // 2:  # Divide by 2 for bidirectional edges\n            # Shuffle and take only what we can handle\n            random.shuffle(potential_base_pairs)\n            potential_base_pairs = potential_base_pairs[:max_additional_edges // 2]\n        \n        # Add the selected base pairs\n        for i, j in potential_base_pairs:\n            edges.append([i, j])\n            edges.append([j, i])  # Bidirectional\n    \n    # Convert edges to tensor\n    edge_index = torch.tensor(edges, dtype=torch.long).t().contiguous()\n    \n    # Get coordinates if available\n    y = None\n    mask = None\n    if labels_df is not None:\n        target_labels = labels_df[labels_df['ID'].str.startswith(target_id + '_')]\n        \n        # Sort by residue ID to match sequence order\n        target_labels = target_labels.sort_values(by='resid')\n        \n        # Check if we have the expected number of residues\n        if len(target_labels) == len(sequence):\n            # Extract coordinates for each residue\n            coordinates = target_labels[['x_1', 'y_1', 'z_1']].values\n            \n            # Create a mask for NaN values (1 for valid, 0 for NaN)\n            valid_mask = ~np.isnan(coordinates).any(axis=1)\n            mask = torch.tensor(valid_mask, dtype=torch.float)\n            \n            # Replace NaN with zeros (we'll mask these during loss calculation)\n            coordinates = np.nan_to_num(coordinates, nan=0.0)\n            \n            y = torch.tensor(coordinates, dtype=torch.float)\n        else:\n            print(f\"Warning: Mismatch in sequence length and label count for {target_id}\")\n    \n    # Create the data object with properly typed target_id (as string)\n    data = Data(x=x, edge_index=edge_index, y=y, mask=mask)\n    \n    # Store target_id as a string attribute\n    data.target_id = str(target_id)\n    \n    return data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:01:44.531570Z","iopub.execute_input":"2025-03-05T09:01:44.531796Z","iopub.status.idle":"2025-03-05T09:01:44.546705Z","shell.execute_reply.started":"2025-03-05T09:01:44.531777Z","shell.execute_reply":"2025-03-05T09:01:44.545934Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_dataset(sequences_df, labels_df=None):\n    dataset = []\n    skipped_count = 0\n    nan_count = 0\n    \n    for idx, row in tqdm(sequences_df.iterrows(), total=len(sequences_df)):\n        target_id = row['target_id']\n        sequence = row['sequence']\n        \n        # Clean sequence - replace any non-standard nucleotides with 'N'\n        # and count how many non-standard nucleotides there are\n        cleaned_sequence = ''\n        non_standard_count = 0\n        \n        for nt in sequence:\n            if nt in nucleotide_mapping:\n                cleaned_sequence += nt\n            else:\n                cleaned_sequence += 'N'  # Placeholder for non-standard nucleotides\n                non_standard_count += 1\n        \n        # If too many non-standard nucleotides (>10%), skip this sequence\n        if non_standard_count / len(sequence) > 0.1:\n            print(f\"Skipping sequence {target_id} with {non_standard_count} non-standard nucleotides\")\n            skipped_count += 1\n            continue\n        \n        # Create graph\n        graph = sequence_to_graph(cleaned_sequence, target_id, labels_df)\n        \n        # Check if we have labels with many NaN values\n        if labels_df is not None and hasattr(graph, 'mask') and graph.mask is not None:\n            nan_percentage = 1.0 - torch.mean(graph.mask).item()\n            if nan_percentage > 0.5:  # If more than 50% coordinates are NaN\n                print(f\"Warning: Sequence {target_id} has {nan_percentage:.1%} NaN coordinates\")\n                nan_count += 1\n        \n        # Add to dataset if no labels needed or valid labels exist\n        if labels_df is None or graph.y is not None:\n            dataset.append(graph)\n    \n    print(f\"Dataset creation: {skipped_count} sequences skipped due to non-standard nucleotides\")\n    print(f\"Dataset creation: {nan_count} sequences have >50% NaN coordinates\")\n    \n    return dataset","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:01:44.547339Z","iopub.execute_input":"2025-03-05T09:01:44.547546Z","iopub.status.idle":"2025-03-05T09:01:44.564513Z","shell.execute_reply.started":"2025-03-05T09:01:44.547528Z","shell.execute_reply":"2025-03-05T09:01:44.563866Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\ndata={\n      \"sequence\":train_sequences['sequence'].to_list(),\n      \"temporal_cutoff\": train_sequences['temporal_cutoff'].to_list(),\n      \"description\": train_sequences['description'].to_list(),\n      \"all_sequences\": train_sequences['all_sequences'].to_list(),\n}\nconfig = {\n    \"cutoff_date\": \"2020-01-01\",\n    \"test_cutoff_date\": \"2022-05-01\",\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:01:48.738940Z","iopub.execute_input":"2025-03-05T09:01:48.739245Z","iopub.status.idle":"2025-03-05T09:01:48.746607Z","shell.execute_reply.started":"2025-03-05T09:01:48.739221Z","shell.execute_reply":"2025-03-05T09:01:48.745769Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Split data into train and test\nall_index = np.arange(len(data['sequence']))\ncutoff_date = pd.Timestamp(config['cutoff_date'])\ntest_cutoff_date = pd.Timestamp(config['test_cutoff_date'])\ntrain_index = [i for i, d in enumerate(data['temporal_cutoff']) if pd.Timestamp(d) <= cutoff_date]\ntest_index = [i for i, d in enumerate(data['temporal_cutoff']) if pd.Timestamp(d) > cutoff_date and pd.Timestamp(d) <= test_cutoff_date]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:01:49.243382Z","iopub.execute_input":"2025-03-05T09:01:49.243650Z","iopub.status.idle":"2025-03-05T09:01:49.250650Z","shell.execute_reply.started":"2025-03-05T09:01:49.243629Z","shell.execute_reply":"2025-03-05T09:01:49.249780Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create training dataset\ntrain_dataset = create_dataset(train_sequences, train_labels)\nprint(f\"Created {len(train_dataset)} graph data objects for training\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:01:51.238505Z","iopub.execute_input":"2025-03-05T09:01:51.238850Z","iopub.status.idle":"2025-03-05T09:02:18.185903Z","shell.execute_reply.started":"2025-03-05T09:01:51.238823Z","shell.execute_reply":"2025-03-05T09:02:18.185106Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_graphs = train_dataset[:len(train_index)]\nval_graphs = train_dataset[:len(train_index)]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:02:18.186933Z","iopub.execute_input":"2025-03-05T09:02:18.187200Z","iopub.status.idle":"2025-03-05T09:02:18.190773Z","shell.execute_reply.started":"2025-03-05T09:02:18.187146Z","shell.execute_reply":"2025-03-05T09:02:18.190037Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T01:15:26.374755Z","iopub.execute_input":"2025-03-05T01:15:26.374961Z","iopub.status.idle":"2025-03-05T01:15:26.393687Z","shell.execute_reply.started":"2025-03-05T01:15:26.374944Z","shell.execute_reply":"2025-03-05T01:15:26.392877Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define the GNN model\nclass RNAStructurePredictor(nn.Module):\n    def __init__(self, input_dim, hidden_dim=256, output_dim=3, num_layers=10, max_seq_len=MAX_SEQ_LEN):\n        super(RNAStructurePredictor, self).__init__()\n        \n        # Initial embedding layer\n        self.embedding = nn.Linear(input_dim, hidden_dim)\n        \n        # SAGEConv layers\n        self.conv_layers = nn.ModuleList()\n        for _ in range(num_layers):\n            self.conv_layers.append(SAGEConv(hidden_dim, hidden_dim))\n        \n        # Output layer for 3D coordinates prediction (x, y, z)\n        self.output = nn.Linear(hidden_dim, output_dim)\n        \n        # Add attention mechanism\n        self.attention = nn.Sequential(\n            nn.Linear(hidden_dim, 1),\n            nn.Sigmoid()\n        )\n        \n        # Position encoding - increase max sequence length\n        self.position_encoder = nn.Embedding(max_seq_len, hidden_dim)\n        \n        # Initialize parameters with Xavier/Glorot\n        for m in self.modules():\n            if isinstance(m, nn.Linear):\n                nn.init.xavier_uniform_(m.weight)\n                if m.bias is not None:\n                    nn.init.zeros_(m.bias)\n        \n    def forward(self, data):\n        x, edge_index = data.x, data.edge_index\n        \n        # Initial embedding\n        x = self.embedding(x)\n        \n        # Add positional information with bounds checking\n        max_pos = self.position_encoder.weight.size(0) - 1  # Maximum allowed index\n        pos = torch.arange(x.size(0), device=x.device)\n        # Clamp position indices to avoid out-of-bounds errors\n        pos = torch.clamp(pos, max=max_pos)\n        x = x + self.position_encoder(pos)\n        \n        # Graph convolution layers\n        for conv in self.conv_layers:\n            x_residual = x\n            x = conv(x, edge_index)\n            x = F.relu(x)\n            x = x + x_residual  # Skip connection\n            x = F.dropout(x, p=0.2, training=self.training)\n        \n        # Apply attention\n        attention_weights = self.attention(x)\n        x = x * attention_weights\n        \n        # Predict 3D coordinates\n        coordinates = self.output(x)\n        \n        return coordinates","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:02:18.192014Z","iopub.execute_input":"2025-03-05T09:02:18.192260Z","iopub.status.idle":"2025-03-05T09:02:18.207544Z","shell.execute_reply.started":"2025-03-05T09:02:18.192239Z","shell.execute_reply":"2025-03-05T09:02:18.206747Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define loss function for 3D coordinate prediction\ndef rmsd_loss(pred, target, mask=None):\n    \"\"\"\n    Root Mean Square Deviation (RMSD) loss function with optional masking for NaN values.\n    Lower RMSD indicates better structural similarity.\n    \n    Args:\n        pred: Predicted coordinates, shape (n_nucleotides, 3)\n        target: Target coordinates, shape (n_nucleotides, 3)\n        mask: Optional mask for valid values, shape (n_nucleotides,)\n    \"\"\"\n    squared_diff = torch.sum((pred - target) ** 2, dim=1)\n    \n    if mask is not None:\n        # Apply mask to consider only valid coordinates\n        # Ensure we don't divide by zero by adding a small epsilon to the sum\n        masked_squared_diff = squared_diff * mask\n        mean_squared_diff = torch.sum(masked_squared_diff) / (torch.sum(mask) + 1e-10)\n    else:\n        mean_squared_diff = torch.mean(squared_diff)\n    \n    rmsd = torch.sqrt(mean_squared_diff)\n    return rmsd","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:02:51.833181Z","iopub.execute_input":"2025-03-05T09:02:51.833497Z","iopub.status.idle":"2025-03-05T09:02:51.838272Z","shell.execute_reply.started":"2025-03-05T09:02:51.833474Z","shell.execute_reply":"2025-03-05T09:02:51.837262Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def calculate_distance_matrix(X,Y,epsilon=1e-4):\n    return (torch.square(X[:,None]-Y[None,:])+epsilon).sum(-1).sqrt()\n\n\ndef dRMSD(pred_x,\n          pred_y,\n          gt_x,\n          gt_y,\n          epsilon=1e-4,Z=10,d_clamp=None):\n    pred_dm=calculate_distance_matrix(pred_x,pred_y)\n    gt_dm=calculate_distance_matrix(gt_x,gt_y)\n\n\n\n    mask=~torch.isnan(gt_dm)\n    mask[torch.eye(mask.shape[0]).bool()]=False\n\n    if d_clamp is not None:\n        rmsd=(torch.square(pred_dm[mask]-gt_dm[mask])+epsilon).clip(0,d_clamp**2)\n    else:\n        rmsd=torch.square(pred_dm[mask]-gt_dm[mask])+epsilon\n\n    return rmsd.sqrt().mean()/Z\n\ndef local_dRMSD(pred_x,\n          pred_y,\n          gt_x,\n          gt_y,\n          epsilon=1e-4,Z=10,d_clamp=30):\n    pred_dm=calculate_distance_matrix(pred_x,pred_y)\n    gt_dm=calculate_distance_matrix(gt_x,gt_y)\n\n\n\n    mask=(~torch.isnan(gt_dm))*(gt_dm<d_clamp)\n    mask[torch.eye(mask.shape[0]).bool()]=False\n\n\n\n    rmsd=torch.square(pred_dm[mask]-gt_dm[mask])+epsilon\n    # rmsd=(torch.square(pred_dm[mask]-gt_dm[mask])+epsilon).sqrt()/Z\n    #rmsd=torch.abs(pred_dm[mask]-gt_dm[mask])/Z\n    return rmsd.sqrt().mean()/Z\n\ndef dRMAE(pred_x,\n          pred_y,\n          gt_x,\n          gt_y,\n          epsilon=1e-4,Z=10,d_clamp=None):\n    pred_dm=calculate_distance_matrix(pred_x,pred_y)\n    gt_dm=calculate_distance_matrix(gt_x,gt_y)\n\n\n\n    mask=~torch.isnan(gt_dm)\n    mask[torch.eye(mask.shape[0]).bool()]=False\n\n    rmsd=torch.abs(pred_dm[mask]-gt_dm[mask])\n\n    return rmsd.mean()/Z\n\nimport torch\n\ndef align_svd_mae(input, target, Z=10):\n    \"\"\"\n    Aligns the input (Nx3) to target (Nx3) using SVD-based Procrustes alignment\n    and computes RMSD loss.\n    \n    Args:\n        input (torch.Tensor): Nx3 tensor representing the input points.\n        target (torch.Tensor): Nx3 tensor representing the target points.\n    \n    Returns:\n        aligned_input (torch.Tensor): Nx3 aligned input.\n        rmsd_loss (torch.Tensor): RMSD loss.\n    \"\"\"\n    assert input.shape == target.shape, \"Input and target must have the same shape\"\n\n    #mask \n    mask=~torch.isnan(target.sum(-1))\n\n    input=input[mask]\n    target=target[mask]\n    \n    # Compute centroids\n    centroid_input = input.mean(dim=0, keepdim=True)\n    centroid_target = target.mean(dim=0, keepdim=True)\n\n    # Center the points\n    input_centered = input - centroid_input.detach()\n    target_centered = target - centroid_target\n\n    # Compute covariance matrix\n    cov_matrix = input_centered.T @ target_centered\n\n    # SVD to find optimal rotation\n    U, S, Vt = torch.svd(cov_matrix)\n\n    # Compute rotation matrix\n    R = Vt @ U.T\n\n    # Ensure a proper rotation (det(R) = 1, no reflection)\n    if torch.det(R) < 0:\n        Vt[-1, :] *= -1\n        R = Vt @ U.T\n\n    # Rotate input\n    aligned_input = (input_centered @ R.T.detach()) + centroid_target.detach()\n\n    # # Compute RMSD loss\n    # rmsd_loss = torch.sqrt(((aligned_input - target) ** 2).mean())\n\n    # rmsd_loss = torch.sqrt(((aligned_input - target) ** 2).mean())\n    \n    # return aligned_input, rmsd_loss\n    return torch.abs(aligned_input-target).mean()/Z","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:02:58.209264Z","iopub.execute_input":"2025-03-05T09:02:58.209555Z","iopub.status.idle":"2025-03-05T09:02:58.220003Z","shell.execute_reply.started":"2025-03-05T09:02:58.209533Z","shell.execute_reply":"2025-03-05T09:02:58.219232Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train(model, train_loader, optimizer, device):\n    model.train()\n    total_loss = 0\n    loss_values = []\n    # Create tqdm progress bar with loss display\n    pbar = tqdm(train_loader, desc='Training')\n    \n    for data in pbar:\n        data = data.to(device)\n        optimizer.zero_grad()\n        \n        # Forward pass\n        pred = model(data)\n        \n        # Calculate loss if labels exist\n        if data.y is not None:\n            # Use mask if available\n            if hasattr(data, 'mask') and data.mask is not None:\n                loss = dRMAE(pred,pred,data.y,data.y) + align_svd_mae(pred, data.y)\n                # loss = rmsd_loss(pred, data.y, data.mask)\n            else:\n                loss = dRMAE(pred,pred,data.y,data.y) + align_svd_mae(pred, data.y)\n                # loss = rmsd_loss(pred, data.y)\n                \n            loss.backward()\n            optimizer.step()\n            total_loss += loss.item()\n            loss_values.append(loss.item())\n            \n            # Update progress bar with current loss\n            pbar.set_postfix({'loss': f'{loss.item():.4f}', 'smooth loss': np.mean(loss_values[-100:])})\n    \n    avg_loss = total_loss / len(train_loader)\n    return avg_loss\n\ndef validate(model, val_loader, device):\n    model.eval()\n    total_loss = 0\n    \n    # Create tqdm progress bar with loss display\n    pbar = tqdm(val_loader, desc='Validation')\n    \n    with torch.no_grad():\n        for data in pbar:\n            data = data.to(device)\n            pred = model(data)\n            \n            if data.y is not None:\n                # Use mask if available\n                if hasattr(data, 'mask') and data.mask is not None:\n                    loss = dRMAE(pred,pred,data.y,data.y) + align_svd_mae(pred, data.y)\n                    # loss = rmsd_loss(pred, data.y, data.mask)\n                else:\n                    loss = dRMAE(pred,pred,data.y,data.y) + align_svd_mae(pred, data.y)\n                    # loss = rmsd_loss(pred, data.y)\n                total_loss += loss.item()\n                \n                # Update progress bar with current loss\n                pbar.set_postfix({'loss': f'{loss.item():.4f}'})\n    \n    avg_loss = total_loss / len(val_loader)\n    return avg_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:03:04.432004Z","iopub.execute_input":"2025-03-05T09:03:04.432333Z","iopub.status.idle":"2025-03-05T09:03:04.440498Z","shell.execute_reply.started":"2025-03-05T09:03:04.432306Z","shell.execute_reply":"2025-03-05T09:03:04.439400Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Function to make predictions on test data\ndef predict(model, test_loader, device):\n    model.eval()\n    predictions = {}\n    \n    with torch.no_grad():\n        for data in test_loader:\n            data = data.to(device)\n            pred = model(data)\n            \n            # Store predictions\n            target_id = data.target_id\n            \n            # If we have ground truth and mask, report metrics\n            if hasattr(data, 'y') and data.y is not None:\n                if hasattr(data, 'mask') and data.mask is not None:\n                    loss = rmsd_loss(pred, data.y, data.mask).item()\n                else:\n                    loss = rmsd_loss(pred, data.y).item()\n                print(f\"Prediction for {target_id}, RMSD: {loss:.4f}\")\n            \n            predictions[target_id] = pred.cpu().numpy()\n    \n    return predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:03:04.618639Z","iopub.execute_input":"2025-03-05T09:03:04.618877Z","iopub.status.idle":"2025-03-05T09:03:04.624004Z","shell.execute_reply.started":"2025-03-05T09:03:04.618857Z","shell.execute_reply":"2025-03-05T09:03:04.623240Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Setup for training\ndevice = 'cpu'\nif torch.cuda.is_available():\n    device = 'cuda'\nelif torch.backends.mps.is_available():\n    device = 'mps'\ndevice = torch.device(device)\nprint(f\"Using device: {device}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:03:20.013925Z","iopub.execute_input":"2025-03-05T09:03:20.014233Z","iopub.status.idle":"2025-03-05T09:03:20.019426Z","shell.execute_reply.started":"2025-03-05T09:03:20.014209Z","shell.execute_reply":"2025-03-05T09:03:20.018491Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create data loaders\ntrain_loader = DataLoader(train_graphs, batch_size=8, shuffle=True, num_workers=4)\nval_loader = DataLoader(val_graphs, batch_size=8, shuffle=False, num_workers=4)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:03:23.853124Z","iopub.execute_input":"2025-03-05T09:03:23.853469Z","iopub.status.idle":"2025-03-05T09:03:23.858881Z","shell.execute_reply.started":"2025-03-05T09:03:23.853443Z","shell.execute_reply":"2025-03-05T09:03:23.857936Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Initialize model\ninput_dim = 4  # One-hot encoding dimension for nucleotides\nmodel = RNAStructurePredictor(input_dim, hidden_dim=1024, output_dim=3, num_layers=15, max_seq_len=MAX_SEQ_LEN).to(device)\nprint(f\"Model initialized with max sequence length of 10000\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:03:35.425975Z","iopub.execute_input":"2025-03-05T09:03:35.426291Z","iopub.status.idle":"2025-03-05T09:03:35.999035Z","shell.execute_reply.started":"2025-03-05T09:03:35.426268Z","shell.execute_reply":"2025-03-05T09:03:35.998268Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"optimizer = torch.optim.Adam(model.parameters(), lr=0.00003)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=12)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:04:14.522773Z","iopub.execute_input":"2025-03-05T09:04:14.523103Z","iopub.status.idle":"2025-03-05T09:04:14.527562Z","shell.execute_reply.started":"2025-03-05T09:04:14.523075Z","shell.execute_reply":"2025-03-05T09:04:14.526647Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Training loop\nnum_epochs = 2000\nbest_val_loss = float('inf')\nearly_stopping_patience = 150\nearly_stopping_counter = 0\n\ntrain_losses = []\nval_losses = []\n\nprint(\"Starting training...\")\nfor epoch in range(num_epochs):\n    # Train\n    train_loss = train(model, train_loader, optimizer, device)\n    train_losses.append(train_loss)\n    \n    # Validate\n    val_loss = validate(model, val_loader, device)\n    val_losses.append(val_loss)\n    \n    # Learning rate scheduler\n    scheduler.step(val_loss)\n    \n    # Early stopping\n    if val_loss < best_val_loss:\n        best_val_loss = val_loss\n        early_stopping_counter = 0\n        # Save best model\n        torch.save(model.state_dict(), \"best_rna_structure_model.pt\")\n    else:\n        early_stopping_counter += 1\n    \n    print(f\"Epoch {epoch+1}/{num_epochs}, \"\n          f\"Train Loss: {train_loss:.4f}, \"\n          f\"Val Loss: {val_loss:.4f}, \"\n          f\"LR: {optimizer.param_groups[0]['lr']:.6f}\")\n    \n    if early_stopping_counter >= early_stopping_patience:\n        print(f\"Early stopping triggered after {epoch+1} epochs\")\n        break\n\nprint(\"Training completed!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T09:04:31.562203Z","iopub.execute_input":"2025-03-05T09:04:31.562517Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Plot training and validation losses\nplt.figure(figsize=(10, 6))\nplt.plot(train_losses, label='Training Loss')\nplt.plot(val_losses, label='Validation Loss')\nplt.xlabel('Epochs')\nplt.ylabel('RMSD Loss')\nplt.title('Training and Validation Losses')\nplt.legend()\nplt.grid(True)\nplt.savefig('training_loss.png')\nplt.close()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T01:16:41.921671Z","iopub.status.idle":"2025-03-05T01:16:41.921975Z","shell.execute_reply":"2025-03-05T01:16:41.921862Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Generate multiple conformations for each RNA sequence\ndef generate_multiple_conformations(model, data, num_conformations=5):\n    \"\"\"\n    Generate multiple structural conformations for an RNA sequence.\n    \n    Args:\n        model: The trained GNN model\n        data: Graph data object containing the RNA sequence\n        num_conformations: Number of conformations to generate (default: 5)\n        \n    Returns:\n        List of numpy arrays, each array has shape (n_nucleotides, 3) for x,y,z coordinates\n    \"\"\"\n    model.eval()\n    conformations = []\n    \n    # Set random seed for reproducibility\n    torch.manual_seed(42)\n    \n    with torch.no_grad():\n        # Generate first conformation (deterministic)\n        base_pred = model(data)\n        base_np = base_pred.cpu().numpy()\n        \n        # Check if base prediction contains NaN values\n        if np.isnan(base_np).any():\n            print(\"Warning: Base prediction contains NaN values. Replacing with zeros.\")\n            base_np = np.nan_to_num(base_np, nan=0.0)\n        \n        # Save the base prediction\n        conformations.append(base_np)\n        \n        # Generate additional conformations with controlled variations\n        for i in range(1, num_conformations):\n            # Use different seeds for different conformations\n            torch.manual_seed(42 + i * 100)  # Larger seed increment for more diversity\n            \n            # Create a copy of the base prediction with a small, controlled variation\n            variation = base_np.copy()\n            \n            # Add random noise with small magnitude (1-5% of the coordinate values)\n            # Calculate standard deviation of base coordinates to scale noise appropriately\n            if not np.all(base_np == 0):  # Check if base_np is not all zeros\n                coord_std = max(np.std(base_np), 0.5)  # Use at least 0.5 to avoid too small noise\n                noise_scale = coord_std * 0.05 * (i + 1)  # Increasing noise for each conformation\n            else:\n                # If base prediction is all zeros (which shouldn't happen normally)\n                noise_scale = 0.5 * (i + 1)\n            \n            # Generate noise and ensure it's not NaN\n            noise = np.random.normal(0, noise_scale, size=variation.shape)\n            \n            # Apply noise to create a new conformation\n            variation += noise\n            \n            # Ensure no NaN values\n            variation = np.nan_to_num(variation, nan=0.0)\n            \n            conformations.append(variation)\n    \n    # Double-check that all conformations are valid and contain no NaNs\n    for i, conf in enumerate(conformations):\n        if np.isnan(conf).any():\n            print(f\"Warning: Conformation {i+1} contains NaN values after processing. Replacing with zeros.\")\n            conformations[i] = np.nan_to_num(conf, nan=0.0)\n    \n    return conformations\n\n# Function to make multiple predictions for test data\ndef predict_multiple_conformations(model, test_loader, device, num_conformations=5):\n    predictions = {}\n    \n    for data in test_loader:\n        data = data.to(device)\n        conformations = generate_multiple_conformations(model, data, num_conformations)\n        \n        # Store predictions - ensure target_id is a hashable type (string)\n        # The target_id could be stored as a list or other non-hashable type\n        if hasattr(data, 'target_id'):\n            # Convert to string if it's not already\n            if isinstance(data.target_id, list) and len(data.target_id) > 0:\n                target_id = str(data.target_id[0])  # Take the first element if it's a list\n            else:\n                target_id = str(data.target_id)  # Convert to string to ensure hashability\n        else:\n            # Generate a unique ID if none exists\n            target_id = f\"unknown_target_{len(predictions)}\"\n            \n        print(f\"Processing target: {target_id}\")\n        predictions[target_id] = conformations\n        \n        # If we have ground truth, report metrics for the first conformation\n        if hasattr(data, 'y') and data.y is not None and len(conformations) > 0:\n            first_conf = torch.tensor(conformations[0], device=device)\n            \n            if hasattr(data, 'mask') and data.mask is not None:\n                loss = rmsd_loss(first_conf, data.y, data.mask).item()\n            else:\n                loss = rmsd_loss(first_conf, data.y).item()\n                \n            print(f\"Prediction for {target_id}, RMSD of first conformation: {loss:.4f}\")\n    \n    return predictions\n\n# Example of how to use the prediction function on test data\ndef process_test_data(test_sequences_path):\n    # Load test sequences\n    test_sequences = pd.read_csv(test_sequences_path)\n    \n    # Create test dataset (without labels)\n    test_dataset = create_dataset(test_sequences)\n    \n    # Create test loader\n    test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)\n    \n    # Make predictions\n    predictions = predict_multiple_conformations(model, test_loader, device)\n    \n    # Format predictions for submission\n    formatted_predictions = []\n    \n    for target_id, conformations in predictions.items():\n        for i, conformation in enumerate(conformations):\n            for j, coords in enumerate(conformation):\n                resid = j + 1  # 1-based indexing\n                row = {\n                    'ID': f\"{target_id}_{resid}\",\n                    f'x_{i+1}': coords[0],\n                    f'y_{i+1}': coords[1],\n                    f'z_{i+1}': coords[2]\n                }\n                formatted_predictions.append(row)\n    \n    # Create submission dataframe\n    submission_df = pd.DataFrame(formatted_predictions)\n    return submission_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T01:15:35.269907Z","iopub.status.idle":"2025-03-05T01:15:35.270263Z","shell.execute_reply":"2025-03-05T01:15:35.270092Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_predictions = process_test_data(\"/kaggle/input/stanford-rna-3d-folding/test_sequences.csv\")\nsub = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/sample_submission.csv\")\nDF_ROWS = []\n\nfor i, row in sub.iterrows():\n    snap = test_predictions[test_predictions['ID'] == row['ID']]\n    x1, y1, z1, x2, y2, z2, x3, y3, z3, x4, y4, z4, x5, y5, z5 = snap['x_1'], snap['y_1'], snap['z_1'], snap['x_2'], snap['y_2'], snap['z_2'], snap['x_3'], snap['y_3'], snap['z_3'], snap['x_4'], snap['y_4'], snap['z_4'], snap['x_5'], snap['y_5'], snap['z_5']\n    x1, y1, z1, x2, y2, z2, x3, y3, z3, x4, y4, z4, x5, y5, z5 = x1.values[0], y1.values[0], z1.values[0], x2.values[1], y2.values[1], z2.values[1], x3.values[2], y3.values[2], z3.values[2], x4.values[3], y4.values[3], z4.values[3], x5.values[4], y5.values[4], z5.values[4]\n    _row = [x1, y1, z1, x2, y2, z2, x3, y3, z3, x4, y4, z4, x5, y5, z5]\n    DF_ROWS.append(_row)\nsub[['x_1', 'y_1', 'z_1', 'x_2', 'y_2', 'z_2', 'x_3', 'y_3', 'z_3', 'x_4', 'y_4', 'z_4', 'x_5', 'y_5', 'z_5']] = DF_ROWS\nsub.head()\nsub.to_csv(\"submission.csv\", index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-05T01:15:35.270978Z","iopub.status.idle":"2025-03-05T01:15:35.271268Z","shell.execute_reply":"2025-03-05T01:15:35.271162Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}