{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":118765,"databundleVersionId":15231210,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":14467782,"sourceType":"datasetVersion","datasetId":9241051}],"dockerImageVersionId":31236,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install /kaggle/input/biopywheel/biopython_wheel/biopython-1.85-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-11T22:06:06.008321Z","iopub.execute_input":"2026-01-11T22:06:06.008891Z","iopub.status.idle":"2026-01-11T22:06:09.589021Z","shell.execute_reply.started":"2026-01-11T22:06:06.008862Z","shell.execute_reply":"2026-01-11T22:06:09.588193Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nfrom scipy.spatial.transform import Rotation as R\nimport random\nfrom Bio import pairwise2\nfrom Bio.Seq import Seq\nimport time\nfrom sklearn.preprocessing import normalize\nfrom scipy.spatial import distance_matrix\n\nimport matplotlib.pyplot as plt\nfrom mpl_toolkits.mplot3d import Axes3D\nimport seaborn as sns\n\nimport warnings\nwarnings.filterwarnings('ignore')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-11T22:06:09.590973Z","iopub.execute_input":"2026-01-11T22:06:09.591295Z","iopub.status.idle":"2026-01-11T22:06:09.596947Z","shell.execute_reply.started":"2026-01-11T22:06:09.591265Z","shell.execute_reply":"2026-01-11T22:06:09.596274Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_seqs = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/train_sequences.csv')\nvalid_seqs = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/validation_sequences.csv')\ntest_seqs = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv')\ntrain_labels = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/train_labels.csv')\nvalid_labels = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/validation_labels.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-11T22:06:09.597839Z","iopub.execute_input":"2026-01-11T22:06:09.598086Z","iopub.status.idle":"2026-01-11T22:06:18.06419Z","shell.execute_reply.started":"2026-01-11T22:06:09.598065Z","shell.execute_reply":"2026-01-11T22:06:18.063407Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"Loaded {len(train_seqs)} training sequences, {len(valid_seqs)} validation sequences, and {len(test_seqs)} test sequences\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-11T22:06:18.06594Z","iopub.execute_input":"2026-01-11T22:06:18.066338Z","iopub.status.idle":"2026-01-11T22:06:18.071484Z","shell.execute_reply.started":"2026-01-11T22:06:18.066311Z","shell.execute_reply":"2026-01-11T22:06:18.070542Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom scipy.stats import entropy\n\ndef plot_enhanced_composition(train_df, valid_df, test_df):\n    \"\"\"\n    Plot base composition distributions, GC content, and sequence complexity.\n    \"\"\"\n    datasets = {'Train': train_df, 'Validation': valid_df, 'Test': test_df}\n    \n    def get_detailed_stats(df, name):\n        stats = []\n        for seq in df['sequence']:\n            seq = seq.upper()\n            length = len(seq)\n            counts = {b: seq.count(b) for b in 'ACGU'}\n            \n            # Basic Ratios\n            gc_content = (counts['G'] + counts['C']) / length * 100\n            \n            # Shannon Entropy: H = -sum(p_i * log2(p_i))\n            probs = [counts[b]/length for b in 'ACGU' if counts[b] > 0]\n            seq_entropy = entropy(probs, base=2)\n            \n            stats.append({\n                'Dataset': name,\n                'GC_Content': gc_content,\n                'Entropy': seq_entropy,\n                **{f'{b}%': (counts[b]/length * 100) for b in 'ACGU'}\n            })\n        return pd.DataFrame(stats)\n\n    # Combine stats for all datasets\n    all_stats = pd.concat([get_detailed_stats(df, name) for name, df in datasets.items()])\n\n    # Plotting\n    fig, axes = plt.subplots(1, 3, figsize=(18, 5))\n\n    # 1. GC Content Distribution (Violin Plot)\n    sns.violinplot(x='Dataset', y='GC_Content', data=all_stats, ax=axes[0], palette='muted')\n    axes[0].set_title('GC Content Distribution')\n    axes[0].set_ylabel('GC %')\n\n    # 2. Base Composition (Mean Bar Chart - updated version of your original)\n    melted_bases = all_stats.melt(id_vars='Dataset', value_vars=['A%', 'C%', 'G%', 'U%'], var_name='Base', value_name='Percentage')\n    sns.barplot(x='Base', y='Percentage', hue='Dataset', data=melted_bases, ax=axes[1])\n    axes[1].set_title('Mean Base Composition')\n\n    # 3. Sequence Complexity (Shannon Entropy)\n    sns.boxplot(x='Dataset', y='Entropy', data=all_stats, ax=axes[2], palette='pastel')\n    axes[2].set_title('Sequence Complexity (Shannon Entropy)')\n    axes[2].set_ylabel('Bits')\n\n    plt.tight_layout()\n    plt.show()\n\n# Usage:\nplot_enhanced_composition(train_seqs, valid_seqs, test_seqs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-11T22:06:18.072665Z","iopub.execute_input":"2026-01-11T22:06:18.073026Z","iopub.status.idle":"2026-01-11T22:06:21.036171Z","shell.execute_reply.started":"2026-01-11T22:06:18.07297Z","shell.execute_reply":"2026-01-11T22:06:21.035409Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Process training labels to create a dictionary mapping target_id to its 3D coordinates\ndef process_labels(labels_df):\n    \"\"\"\n    Process labels dataframe to create a dictionary mapping target_id to coordinates\n    \"\"\"\n    coords_dict = {}\n    \n    # Group by target ID\n    for id_prefix, group in labels_df.groupby(lambda x: labels_df['ID'][x].rsplit('_', 1)[0]):\n        # Extract just the coordinates columns for the first structure (x_1, y_1, z_1)\n        coords = []\n        for _, row in group.sort_values('resid').iterrows():\n            coords.append([row['x_1'], row['y_1'], row['z_1']])\n        \n        coords_dict[id_prefix] = np.array(coords)\n    \n    return coords_dict\n\n# Process training labels\nprint(\"Processing training labels...\")\ntrain_coords_dict = process_labels(train_labels)\nvalid_coords_dict = process_labels(valid_labels)\nprint(f\"Processed coordinates for {len(train_coords_dict)} training structures and {len(valid_coords_dict)} validation structures\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-11T22:06:21.037226Z","iopub.execute_input":"2026-01-11T22:06:21.037534Z","iopub.status.idle":"2026-01-11T22:13:04.507778Z","shell.execute_reply.started":"2026-01-11T22:06:21.0375Z","shell.execute_reply":"2026-01-11T22:13:04.507061Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\nimport random\n\ndef plot_enhanced_structure_stats(coords_dict, sample_size=100):\n    sequential_dists = []\n    all_pairwise_dists = []\n    radii_of_gyration = []\n    clash_count = 0\n    total_pairs_checked = 0\n    \n    sampled_ids = random.sample(list(coords_dict.keys()), min(sample_size, len(coords_dict)))\n    \n    for target_id in sampled_ids:\n        coords = coords_dict[target_id]\n        \n        # 1. Radius of Gyration (Rg) - Measures compactness\n        # Formula: sqrt(mean(squared distance from centroid))\n        centroid = np.mean(coords, axis=0)\n        rg = np.sqrt(np.mean(np.sum((coords - centroid)**2, axis=1)))\n        radii_of_gyration.append(rg)\n        \n        # 2. Sequential distances\n        for i in range(len(coords) - 1):\n            dist = np.linalg.norm(coords[i+1] - coords[i])\n            sequential_dists.append(dist)\n        \n        # 3. Pairwise distances & Physical Clashes\n        if len(coords) > 2:\n            # We sample pairs to keep it fast\n            idx = np.random.choice(len(coords), min(50, len(coords)), replace=False)\n            for i in range(len(idx)):\n                for j in range(i + 1, len(idx)):\n                    dist = np.linalg.norm(coords[idx[i]] - coords[idx[j]])\n                    all_pairwise_dists.append(dist)\n                    total_pairs_checked += 1\n                    if dist < 3.0:  # Physical clash threshold for C1' atoms\n                        clash_count += 1\n\n    # Visualization\n    fig, axes = plt.subplots(1, 3, figsize=(20, 5))\n    \n    # Plot 1: Sequential Distances (Backbone integrity)\n    axes[0].hist(sequential_dists, bins=50, color='steelblue', alpha=0.7)\n    median_seq = np.median(sequential_dists)\n    axes[0].axvline(median_seq, color='red', linestyle='--', label=f'Med: {median_seq:.2f}Å')\n    axes[0].set_title('Sequential C1\\' Distances\\n(Backbone Integrity)')\n    axes[0].set_xlabel('Distance (Å)')\n    axes[0].legend()\n\n    # Plot 2: Pairwise Distances (Global Scale)\n    axes[1].hist(all_pairwise_dists, bins=50, color='coral', alpha=0.7)\n    axes[1].set_title('Pairwise Distances\\n(Overall Scale)')\n    axes[1].set_xlabel('Distance (Å)')\n    clash_pct = (clash_count / total_pairs_checked) * 100 if total_pairs_checked > 0 else 0\n    axes[1].annotate(f'Clash Rate (<3Å): {clash_pct:.2f}%', xy=(0.5, 0.9), \n                     xycoords='axes fraction', color='red', weight='bold')\n\n    # Plot 3: Radius of Gyration (Compactness)\n    axes[2].hist(radii_of_gyration, bins=30, color='seagreen', alpha=0.7)\n    axes[2].set_title('Radius of Gyration\\n(Structure Compactness)')\n    axes[2].set_xlabel('Rg (Å)')\n    \n    plt.tight_layout()\n    plt.show()\n\n    # Print Descriptive Summary\n    print(f\"--- Structural Statistics Summary ({len(sampled_ids)} samples) ---\")\n    print(f\"Sequential Dist: Mean={np.mean(sequential_dists):.2f}, Std={np.std(sequential_dists):.2f}\")\n    print(f\"Global Scale:    Max Pairwise={np.max(all_pairwise_dists):.2f}, Median Rg={np.median(radii_of_gyration):.2f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-11T22:13:04.508682Z","iopub.execute_input":"2026-01-11T22:13:04.508966Z","iopub.status.idle":"2026-01-11T22:13:04.520694Z","shell.execute_reply.started":"2026-01-11T22:13:04.508943Z","shell.execute_reply":"2026-01-11T22:13:04.520025Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#plot_enhanced_structure_stats(train_coords_dict, sample_size=100)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-11T22:13:04.521526Z","iopub.execute_input":"2026-01-11T22:13:04.521734Z","iopub.status.idle":"2026-01-11T22:13:06.483952Z","shell.execute_reply.started":"2026-01-11T22:13:04.521715Z","shell.execute_reply":"2026-01-11T22:13:06.483293Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-11T22:43:17.229408Z","iopub.execute_input":"2026-01-11T22:43:17.229877Z","iopub.status.idle":"2026-01-11T22:43:17.612874Z","shell.execute_reply.started":"2026-01-11T22:43:17.229842Z","shell.execute_reply":"2026-01-11T22:43:17.611768Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nRNA 3D Structure Prediction with Graph Diffusion Model\n=======================================================\n\nA Graph Neural Network-based diffusion model for RNA structure prediction.\nUses PyTorch (no torch_scatter dependency) with:\n- E(3) Equivariant Graph Neural Networks (EGNN)\n- Graph-based message passing for coordinate denoising\n- Edge features encoding spatial and sequential relationships\n- Score-based diffusion on graphs\n- DDPM/DDIM sampling with graph structure\n\nKey innovations over standard diffusion:\n1. Graph representation of RNA with sequential + k-NN spatial edges\n2. E(3) equivariant coordinate updates\n3. Edge-conditioned message passing\n4. Multi-scale graph convolutions\n\nKaggle Competition: Stanford RNA 3D Folding 2\n\"\"\"\n\nimport numpy as np\nimport pandas as pd\nimport time\nimport math\nfrom typing import List, Tuple, Optional, Dict, Any\nfrom dataclasses import dataclass\nfrom collections import defaultdict\nfrom tqdm import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingLR, OneCycleLR\n\n# Set device\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {DEVICE}\")\n\n\n# =============================================================================\n# SCATTER OPERATIONS (Native PyTorch Implementation)\n# =============================================================================\n\ndef scatter_add(src: torch.Tensor, index: torch.Tensor, dim: int = 0,\n                dim_size: int = None) -> torch.Tensor:\n    \"\"\"\n    Scatter add operation using native PyTorch.\n    \n    Args:\n        src: Source tensor of shape (E, ...) \n        index: Index tensor of shape (E,)\n        dim: Dimension along which to scatter\n        dim_size: Size of output dimension\n        \n    Returns:\n        Output tensor with scattered values summed\n    \"\"\"\n    if dim_size is None:\n        dim_size = index.max().item() + 1\n    \n    # Create output tensor\n    shape = list(src.shape)\n    shape[dim] = dim_size\n    out = torch.zeros(shape, dtype=src.dtype, device=src.device)\n    \n    # Expand index to match src dimensions\n    index_expanded = index.view(-1, *([1] * (src.dim() - 1)))\n    index_expanded = index_expanded.expand_as(src)\n    \n    # Scatter add\n    out.scatter_add_(dim, index_expanded, src)\n    \n    return out\n\n\ndef scatter_mean(src: torch.Tensor, index: torch.Tensor, dim: int = 0,\n                 dim_size: int = None) -> torch.Tensor:\n    \"\"\"\n    Scatter mean operation using native PyTorch.\n    \"\"\"\n    if dim_size is None:\n        dim_size = index.max().item() + 1\n    \n    # Sum\n    out_sum = scatter_add(src, index, dim, dim_size)\n    \n    # Count\n    ones = torch.ones(index.shape[0], dtype=src.dtype, device=src.device)\n    count = scatter_add(ones, index, dim=0, dim_size=dim_size)\n    count = count.clamp(min=1)\n    \n    # Reshape count for broadcasting\n    shape = [1] * out_sum.dim()\n    shape[dim] = dim_size\n    count = count.view(*shape)\n    \n    return out_sum / count\n\n\n# =============================================================================\n# DATA LOADING AND PREPROCESSING\n# =============================================================================\n\ndef load_competition_data(data_dir: str = '/kaggle/input/stanford-rna-3d-folding-2'):\n    \"\"\"Load all competition data files.\"\"\"\n    print(\"Loading competition data...\")\n    \n    train_seqs = pd.read_csv(f'{data_dir}/train_sequences.csv')\n    valid_seqs = pd.read_csv(f'{data_dir}/validation_sequences.csv')\n    test_seqs = pd.read_csv(f'{data_dir}/test_sequences.csv')\n    train_labels = pd.read_csv(f'{data_dir}/train_labels.csv')\n    valid_labels = pd.read_csv(f'{data_dir}/validation_labels.csv')\n    \n    print(f\"  Train sequences: {len(train_seqs)}\")\n    print(f\"  Valid sequences: {len(valid_seqs)}\")\n    print(f\"  Test sequences: {len(test_seqs)}\")\n    print(f\"  Train labels: {len(train_labels)}\")\n    print(f\"  Valid labels: {len(valid_labels)}\")\n    \n    return train_seqs, valid_seqs, test_seqs, train_labels, valid_labels\n\n\ndef build_coords_dict(labels_df: pd.DataFrame) -> Dict[str, np.ndarray]:\n    \"\"\"Build dictionary mapping target_id to 3D coordinates.\"\"\"\n    coords_dict = {}\n    \n    labels_df = labels_df.copy()\n    labels_df['target_id'] = labels_df['ID'].apply(lambda x: '_'.join(x.split('_')[:-1]))\n    \n    for target_id, group in labels_df.groupby('target_id'):\n        group = group.sort_values('resid')\n        coords = group[['x_1', 'y_1', 'z_1']].values.astype(np.float32)\n        \n        if not np.isnan(coords).any():\n            coords_dict[target_id] = coords\n    \n    return coords_dict\n\n\n# =============================================================================\n# GRAPH CONSTRUCTION\n# =============================================================================\n\nclass RNAGraphBuilder:\n    \"\"\"Builds graph representation of RNA molecules.\"\"\"\n    \n    NUCLEOTIDE_MAP = {'A': 0, 'C': 1, 'G': 2, 'U': 3, 'T': 3, 'N': 4}\n    BASE_PAIRS = {('A', 'U'), ('U', 'A'), ('G', 'C'), ('C', 'G'), ('G', 'U'), ('U', 'G')}\n    \n    def __init__(self, k_neighbors: int = 10, max_seq_dist: int = 5, use_base_pairing: bool = True):\n        self.k_neighbors = k_neighbors\n        self.max_seq_dist = max_seq_dist\n        self.use_base_pairing = use_base_pairing\n    \n    def encode_sequence(self, sequence: str) -> torch.Tensor:\n        \"\"\"One-hot encode nucleotide sequence.\"\"\"\n        n = len(sequence)\n        encoded = torch.zeros(n, 5)\n        for i, nt in enumerate(sequence.upper()):\n            idx = self.NUCLEOTIDE_MAP.get(nt, 4)\n            encoded[i, idx] = 1.0\n        return encoded\n    \n    def get_positional_encoding(self, n: int, d_model: int = 32) -> torch.Tensor:\n        \"\"\"Sinusoidal positional encoding for sequence positions.\"\"\"\n        position = torch.arange(n).unsqueeze(1).float()\n        div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))\n        \n        pe = torch.zeros(n, d_model)\n        pe[:, 0::2] = torch.sin(position * div_term)\n        pe[:, 1::2] = torch.cos(position * div_term)\n        \n        return pe\n    \n    def build_sequential_edges(self, n: int) -> Tuple[torch.Tensor, torch.Tensor]:\n        \"\"\"Build sequential backbone edges with edge types.\"\"\"\n        edges = []\n        edge_types = []\n        \n        for i in range(n):\n            for offset in range(1, self.max_seq_dist + 1):\n                if i + offset < n:\n                    edges.append([i, i + offset])\n                    edge_types.append(offset)\n                    edges.append([i + offset, i])\n                    edge_types.append(offset)\n        \n        if len(edges) == 0:\n            return torch.zeros(2, 0, dtype=torch.long), torch.zeros(0, dtype=torch.long)\n        \n        edge_index = torch.tensor(edges, dtype=torch.long).t()\n        edge_types = torch.tensor(edge_types, dtype=torch.long)\n        \n        return edge_index, edge_types\n    \n    def build_knn_edges(self, coords: torch.Tensor, k: int = None) -> torch.Tensor:\n        \"\"\"Build k-NN spatial edges based on 3D coordinates.\"\"\"\n        if k is None:\n            k = self.k_neighbors\n        \n        n = coords.shape[0]\n        k = min(k, n - 1)\n        \n        if k <= 0:\n            return torch.zeros(2, 0, dtype=torch.long)\n        \n        diff = coords.unsqueeze(0) - coords.unsqueeze(1)\n        dist = torch.sqrt((diff ** 2).sum(dim=-1) + 1e-8)\n        dist.fill_diagonal_(float('inf'))\n        \n        _, indices = dist.topk(k, dim=1, largest=False)\n        \n        src = torch.arange(n).unsqueeze(1).expand(-1, k).flatten()\n        dst = indices.flatten()\n        \n        edge_index = torch.stack([src, dst], dim=0)\n        \n        return edge_index\n    \n    def build_base_pair_edges(self, sequence: str, coords: torch.Tensor = None,\n                              distance_threshold: float = 15.0) -> torch.Tensor:\n        \"\"\"Build potential base-pairing edges.\"\"\"\n        n = len(sequence)\n        edges = []\n        sequence = sequence.upper()\n        \n        for i in range(n):\n            for j in range(i + 4, n):\n                nt_i, nt_j = sequence[i], sequence[j]\n                \n                if (nt_i, nt_j) in self.BASE_PAIRS:\n                    if coords is not None:\n                        dist = torch.norm(coords[i] - coords[j]).item()\n                        if dist > distance_threshold:\n                            continue\n                    \n                    edges.append([i, j])\n                    edges.append([j, i])\n        \n        if len(edges) == 0:\n            return torch.zeros(2, 0, dtype=torch.long)\n        \n        return torch.tensor(edges, dtype=torch.long).t()\n    \n    def build_graph(self, sequence: str, coords: torch.Tensor = None) -> Dict[str, torch.Tensor]:\n        \"\"\"Build complete graph representation.\"\"\"\n        n = len(sequence)\n        \n        seq_encoding = self.encode_sequence(sequence)\n        pos_encoding = self.get_positional_encoding(n, d_model=32)\n        node_features = torch.cat([seq_encoding, pos_encoding], dim=-1)\n        \n        seq_edges, seq_types = self.build_sequential_edges(n)\n        \n        all_edges = [seq_edges]\n        all_types = [seq_types]\n        \n        if coords is not None:\n            knn_edges = self.build_knn_edges(coords)\n            knn_types = torch.full((knn_edges.shape[1],), self.max_seq_dist + 1, dtype=torch.long)\n            all_edges.append(knn_edges)\n            all_types.append(knn_types)\n            \n            if self.use_base_pairing:\n                bp_edges = self.build_base_pair_edges(sequence, coords)\n                bp_types = torch.full((bp_edges.shape[1],), self.max_seq_dist + 2, dtype=torch.long)\n                all_edges.append(bp_edges)\n                all_types.append(bp_types)\n        \n        edge_index = torch.cat(all_edges, dim=1)\n        edge_type = torch.cat(all_types, dim=0)\n        \n        if edge_index.shape[1] > 0:\n            edge_index, unique_idx = torch.unique(edge_index, dim=1, return_inverse=True)\n            edge_type_new = torch.zeros(edge_index.shape[1], dtype=torch.long)\n            for i, idx in enumerate(unique_idx):\n                edge_type_new[idx] = edge_type[i]\n            edge_type = edge_type_new\n        \n        return {\n            'node_features': node_features,\n            'edge_index': edge_index,\n            'edge_type': edge_type,\n            'num_nodes': n\n        }\n\n\n# =============================================================================\n# DATASET\n# =============================================================================\n\nclass RNAGraphDataset(Dataset):\n    \"\"\"Dataset for RNA graphs with 3D structures.\"\"\"\n    \n    def __init__(self, sequences_df: pd.DataFrame, coords_dict: Dict[str, np.ndarray],\n                 max_len: int = 512):\n        self.max_len = max_len\n        self.graph_builder = RNAGraphBuilder()\n        self.samples = []\n        \n        seq_lookup = dict(zip(sequences_df['target_id'], sequences_df['sequence']))\n        \n        for target_id, coords in coords_dict.items():\n            if target_id in seq_lookup:\n                seq = seq_lookup[target_id]\n                if len(seq) <= max_len and len(seq) == len(coords):\n                    self.samples.append({\n                        'target_id': target_id,\n                        'sequence': seq,\n                        'coords': coords\n                    })\n        \n        print(f\"Dataset: {len(self.samples)} valid samples\")\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n        \n        coords = torch.tensor(sample['coords'], dtype=torch.float32)\n        coords = coords - coords.mean(dim=0, keepdim=True)\n        \n        graph = self.graph_builder.build_graph(sample['sequence'], coords)\n        \n        return {\n            'node_features': graph['node_features'],\n            'edge_index': graph['edge_index'],\n            'edge_type': graph['edge_type'],\n            'coords': coords,\n            'length': len(sample['sequence']),\n            'sequence': sample['sequence'],\n            'target_id': sample['target_id']\n        }\n\n\ndef collate_graph_batch(batch: List[Dict]) -> Dict[str, torch.Tensor]:\n    \"\"\"Collate function for batching graphs.\"\"\"\n    node_features_list = []\n    coords_list = []\n    edge_index_list = []\n    edge_type_list = []\n    batch_idx_list = []\n    lengths = []\n    sequences = []\n    \n    node_offset = 0\n    \n    for i, item in enumerate(batch):\n        n = item['length']\n        \n        node_features_list.append(item['node_features'])\n        coords_list.append(item['coords'])\n        \n        edges = item['edge_index'].clone()\n        if edges.shape[1] > 0:\n            edges += node_offset\n        edge_index_list.append(edges)\n        edge_type_list.append(item['edge_type'])\n        \n        batch_idx_list.append(torch.full((n,), i, dtype=torch.long))\n        \n        lengths.append(n)\n        sequences.append(item['sequence'])\n        \n        node_offset += n\n    \n    return {\n        'node_features': torch.cat(node_features_list, dim=0),\n        'coords': torch.cat(coords_list, dim=0),\n        'edge_index': torch.cat(edge_index_list, dim=1),\n        'edge_type': torch.cat(edge_type_list, dim=0),\n        'batch': torch.cat(batch_idx_list, dim=0),\n        'lengths': torch.tensor(lengths),\n        'sequences': sequences,\n        'ptr': torch.tensor([0] + list(np.cumsum(lengths)))\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-11T23:02:04.761444Z","iopub.execute_input":"2026-01-11T23:02:04.762502Z","iopub.status.idle":"2026-01-11T23:02:04.805428Z","shell.execute_reply.started":"2026-01-11T23:02:04.762467Z","shell.execute_reply":"2026-01-11T23:02:04.804599Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nRNA 3D Structure Prediction with Graph Diffusion Model\n=======================================================\n\nA Graph Neural Network-based diffusion model for RNA structure prediction.\nUses PyTorch (no torch_scatter dependency) with:\n- E(3) Equivariant Graph Neural Networks (EGNN)\n- Graph-based message passing for coordinate denoising\n- Edge features encoding spatial and sequential relationships\n- Score-based diffusion on graphs\n- DDPM/DDIM sampling with graph structure\n\nKaggle Competition: Stanford RNA 3D Folding 2\n\"\"\"\n\nimport numpy as np\nimport pandas as pd\nimport time\nimport math\nfrom typing import List, Tuple, Optional, Dict, Any\nfrom collections import defaultdict\nfrom tqdm import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import OneCycleLR\n\n# Set device - with error handling\ndef get_device():\n    if torch.cuda.is_available():\n        try:\n            # Test CUDA availability\n            torch.cuda.init()\n            _ = torch.zeros(1).cuda()\n            return torch.device('cuda')\n        except Exception as e:\n            print(f\"CUDA initialization failed: {e}\")\n            return torch.device('cpu')\n    return torch.device('cpu')\n\nDEVICE = get_device()\nprint(f\"Using device: {DEVICE}\")\n\n\n# =============================================================================\n# SCATTER OPERATIONS (Native PyTorch Implementation)\n# =============================================================================\n\ndef scatter_add(src: torch.Tensor, index: torch.Tensor, dim: int = 0,\n                dim_size: int = None) -> torch.Tensor:\n    \"\"\"Scatter add operation using native PyTorch.\"\"\"\n    if index.numel() == 0:\n        if dim_size is None:\n            dim_size = 1\n        shape = list(src.shape)\n        shape[dim] = dim_size\n        return torch.zeros(shape, dtype=src.dtype, device=src.device)\n    \n    if dim_size is None:\n        dim_size = int(index.max().item()) + 1\n    \n    shape = list(src.shape)\n    shape[dim] = dim_size\n    out = torch.zeros(shape, dtype=src.dtype, device=src.device)\n    \n    # Expand index to match src dimensions\n    expand_shape = [-1] + [1] * (src.dim() - 1)\n    index_expanded = index.view(*expand_shape).expand_as(src)\n    \n    out.scatter_add_(dim, index_expanded, src)\n    return out\n\n\ndef scatter_mean(src: torch.Tensor, index: torch.Tensor, dim: int = 0,\n                 dim_size: int = None) -> torch.Tensor:\n    \"\"\"Scatter mean operation using native PyTorch.\"\"\"\n    if index.numel() == 0:\n        if dim_size is None:\n            dim_size = 1\n        shape = list(src.shape)\n        shape[dim] = dim_size\n        return torch.zeros(shape, dtype=src.dtype, device=src.device)\n    \n    if dim_size is None:\n        dim_size = int(index.max().item()) + 1\n    \n    out_sum = scatter_add(src, index, dim, dim_size)\n    \n    ones = torch.ones(index.shape[0], dtype=src.dtype, device=src.device)\n    count = scatter_add(ones, index, dim=0, dim_size=dim_size)\n    count = count.clamp(min=1)\n    \n    # Reshape count for broadcasting\n    view_shape = [1] * out_sum.dim()\n    view_shape[dim] = dim_size\n    count = count.view(*view_shape)\n    \n    return out_sum / count\n\n\n# =============================================================================\n# DATA LOADING AND PREPROCESSING\n# =============================================================================\n\ndef load_competition_data(data_dir: str = '/kaggle/input/stanford-rna-3d-folding-2'):\n    \"\"\"Load all competition data files.\"\"\"\n    print(\"Loading competition data...\")\n    \n    train_seqs = pd.read_csv(f'{data_dir}/train_sequences.csv')\n    valid_seqs = pd.read_csv(f'{data_dir}/validation_sequences.csv')\n    test_seqs = pd.read_csv(f'{data_dir}/test_sequences.csv')\n    train_labels = pd.read_csv(f'{data_dir}/train_labels.csv')\n    valid_labels = pd.read_csv(f'{data_dir}/validation_labels.csv')\n    \n    print(f\"  Train sequences: {len(train_seqs)}\")\n    print(f\"  Valid sequences: {len(valid_seqs)}\")\n    print(f\"  Test sequences: {len(test_seqs)}\")\n    print(f\"  Train labels: {len(train_labels)}\")\n    print(f\"  Valid labels: {len(valid_labels)}\")\n    \n    return train_seqs, valid_seqs, test_seqs, train_labels, valid_labels\n\n\ndef build_coords_dict(labels_df: pd.DataFrame) -> Dict[str, np.ndarray]:\n    \"\"\"Build dictionary mapping target_id to 3D coordinates.\"\"\"\n    coords_dict = {}\n    \n    labels_df = labels_df.copy()\n    labels_df['target_id'] = labels_df['ID'].apply(lambda x: '_'.join(x.split('_')[:-1]))\n    \n    for target_id, group in labels_df.groupby('target_id'):\n        group = group.sort_values('resid')\n        coords = group[['x_1', 'y_1', 'z_1']].values.astype(np.float32)\n        \n        if not np.isnan(coords).any():\n            coords_dict[target_id] = coords\n    \n    return coords_dict\n\n\n# =============================================================================\n# GRAPH CONSTRUCTION\n# =============================================================================\n\nclass RNAGraphBuilder:\n    \"\"\"Builds graph representation of RNA molecules.\"\"\"\n    \n    NUCLEOTIDE_MAP = {'A': 0, 'C': 1, 'G': 2, 'U': 3, 'T': 3, 'N': 4}\n    BASE_PAIRS = {('A', 'U'), ('U', 'A'), ('G', 'C'), ('C', 'G'), ('G', 'U'), ('U', 'G')}\n    \n    def __init__(self, k_neighbors: int = 10, max_seq_dist: int = 5, use_base_pairing: bool = True):\n        self.k_neighbors = k_neighbors\n        self.max_seq_dist = max_seq_dist\n        self.use_base_pairing = use_base_pairing\n    \n    def encode_sequence(self, sequence: str) -> torch.Tensor:\n        \"\"\"One-hot encode nucleotide sequence.\"\"\"\n        n = len(sequence)\n        encoded = torch.zeros(n, 5, dtype=torch.float32)\n        for i, nt in enumerate(sequence.upper()):\n            idx = self.NUCLEOTIDE_MAP.get(nt, 4)\n            encoded[i, idx] = 1.0\n        return encoded\n    \n    def get_positional_encoding(self, n: int, d_model: int = 32) -> torch.Tensor:\n        \"\"\"Sinusoidal positional encoding for sequence positions.\"\"\"\n        position = torch.arange(n, dtype=torch.float32).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, d_model, 2, dtype=torch.float32) * (-math.log(10000.0) / d_model))\n        \n        pe = torch.zeros(n, d_model, dtype=torch.float32)\n        pe[:, 0::2] = torch.sin(position * div_term)\n        pe[:, 1::2] = torch.cos(position * div_term)\n        \n        return pe\n    \n    def build_sequential_edges(self, n: int) -> Tuple[torch.Tensor, torch.Tensor]:\n        \"\"\"Build sequential backbone edges with edge types.\"\"\"\n        edges = []\n        edge_types = []\n        \n        for i in range(n):\n            for offset in range(1, self.max_seq_dist + 1):\n                if i + offset < n:\n                    edges.append([i, i + offset])\n                    edge_types.append(offset)\n                    edges.append([i + offset, i])\n                    edge_types.append(offset)\n        \n        if len(edges) == 0:\n            return torch.zeros(2, 0, dtype=torch.long), torch.zeros(0, dtype=torch.long)\n        \n        edge_index = torch.tensor(edges, dtype=torch.long).t().contiguous()\n        edge_types = torch.tensor(edge_types, dtype=torch.long)\n        \n        return edge_index, edge_types\n    \n    def build_knn_edges(self, coords: torch.Tensor, k: int = None) -> torch.Tensor:\n        \"\"\"Build k-NN spatial edges based on 3D coordinates.\"\"\"\n        if k is None:\n            k = self.k_neighbors\n        \n        n = coords.shape[0]\n        k = min(k, n - 1)\n        \n        if k <= 0:\n            return torch.zeros(2, 0, dtype=torch.long)\n        \n        # Compute pairwise distances\n        diff = coords.unsqueeze(0) - coords.unsqueeze(1)\n        dist = torch.sqrt((diff ** 2).sum(dim=-1) + 1e-8)\n        dist.fill_diagonal_(float('inf'))\n        \n        _, indices = dist.topk(k, dim=1, largest=False)\n        \n        src = torch.arange(n).unsqueeze(1).expand(-1, k).flatten()\n        dst = indices.flatten()\n        \n        edge_index = torch.stack([src, dst], dim=0).contiguous()\n        \n        return edge_index\n    \n    def build_base_pair_edges(self, sequence: str, coords: torch.Tensor = None,\n                              distance_threshold: float = 15.0) -> torch.Tensor:\n        \"\"\"Build potential base-pairing edges.\"\"\"\n        n = len(sequence)\n        edges = []\n        sequence = sequence.upper()\n        \n        for i in range(n):\n            for j in range(i + 4, n):\n                nt_i, nt_j = sequence[i], sequence[j]\n                \n                if (nt_i, nt_j) in self.BASE_PAIRS:\n                    if coords is not None:\n                        dist = torch.norm(coords[i] - coords[j]).item()\n                        if dist > distance_threshold:\n                            continue\n                    \n                    edges.append([i, j])\n                    edges.append([j, i])\n        \n        if len(edges) == 0:\n            return torch.zeros(2, 0, dtype=torch.long)\n        \n        return torch.tensor(edges, dtype=torch.long).t().contiguous()\n    \n    def build_graph(self, sequence: str, coords: torch.Tensor = None) -> Dict[str, torch.Tensor]:\n        \"\"\"Build complete graph representation.\"\"\"\n        n = len(sequence)\n        \n        seq_encoding = self.encode_sequence(sequence)\n        pos_encoding = self.get_positional_encoding(n, d_model=32)\n        node_features = torch.cat([seq_encoding, pos_encoding], dim=-1)\n        \n        seq_edges, seq_types = self.build_sequential_edges(n)\n        \n        all_edges = [seq_edges]\n        all_types = [seq_types]\n        \n        if coords is not None:\n            knn_edges = self.build_knn_edges(coords)\n            knn_types = torch.full((knn_edges.shape[1],), self.max_seq_dist + 1, dtype=torch.long)\n            all_edges.append(knn_edges)\n            all_types.append(knn_types)\n            \n            if self.use_base_pairing:\n                bp_edges = self.build_base_pair_edges(sequence, coords)\n                bp_types = torch.full((bp_edges.shape[1],), self.max_seq_dist + 2, dtype=torch.long)\n                all_edges.append(bp_edges)\n                all_types.append(bp_types)\n        \n        edge_index = torch.cat(all_edges, dim=1)\n        edge_type = torch.cat(all_types, dim=0)\n        \n        # Remove duplicate edges\n        if edge_index.shape[1] > 0:\n            # Create unique edge identifiers\n            edge_hash = edge_index[0] * n + edge_index[1]\n            _, unique_indices = torch.unique(edge_hash, return_inverse=True)\n            \n            # Get unique edges\n            seen = set()\n            keep_mask = []\n            for i in range(edge_index.shape[1]):\n                h = edge_hash[i].item()\n                if h not in seen:\n                    seen.add(h)\n                    keep_mask.append(i)\n            \n            if len(keep_mask) > 0:\n                keep_mask = torch.tensor(keep_mask, dtype=torch.long)\n                edge_index = edge_index[:, keep_mask].contiguous()\n                edge_type = edge_type[keep_mask]\n        \n        return {\n            'node_features': node_features,\n            'edge_index': edge_index,\n            'edge_type': edge_type,\n            'num_nodes': n\n        }\n\n\n# =============================================================================\n# DATASET\n# =============================================================================\n\nclass RNAGraphDataset(Dataset):\n    \"\"\"Dataset for RNA graphs with 3D structures.\"\"\"\n    \n    def __init__(self, sequences_df: pd.DataFrame, coords_dict: Dict[str, np.ndarray],\n                 max_len: int = 512):\n        self.max_len = max_len\n        self.graph_builder = RNAGraphBuilder()\n        self.samples = []\n        \n        seq_lookup = dict(zip(sequences_df['target_id'], sequences_df['sequence']))\n        \n        for target_id, coords in coords_dict.items():\n            if target_id in seq_lookup:\n                seq = seq_lookup[target_id]\n                if len(seq) <= max_len and len(seq) == len(coords):\n                    self.samples.append({\n                        'target_id': target_id,\n                        'sequence': seq,\n                        'coords': coords\n                    })\n        \n        print(f\"Dataset: {len(self.samples)} valid samples\")\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n        \n        coords = torch.tensor(sample['coords'], dtype=torch.float32)\n        coords = coords - coords.mean(dim=0, keepdim=True)\n        \n        graph = self.graph_builder.build_graph(sample['sequence'], coords)\n        \n        return {\n            'node_features': graph['node_features'],\n            'edge_index': graph['edge_index'],\n            'edge_type': graph['edge_type'],\n            'coords': coords,\n            'length': len(sample['sequence']),\n            'sequence': sample['sequence'],\n            'target_id': sample['target_id']\n        }\n\n\ndef collate_graph_batch(batch: List[Dict]) -> Dict[str, torch.Tensor]:\n    \"\"\"Collate function for batching graphs.\"\"\"\n    node_features_list = []\n    coords_list = []\n    edge_index_list = []\n    edge_type_list = []\n    batch_idx_list = []\n    lengths = []\n    sequences = []\n    \n    node_offset = 0\n    \n    for i, item in enumerate(batch):\n        n = item['length']\n        \n        node_features_list.append(item['node_features'])\n        coords_list.append(item['coords'])\n        \n        edges = item['edge_index'].clone()\n        if edges.numel() > 0:\n            edges = edges + node_offset\n        edge_index_list.append(edges)\n        edge_type_list.append(item['edge_type'])\n        \n        batch_idx_list.append(torch.full((n,), i, dtype=torch.long))\n        \n        lengths.append(n)\n        sequences.append(item['sequence'])\n        \n        node_offset += n\n    \n    return {\n        'node_features': torch.cat(node_features_list, dim=0),\n        'coords': torch.cat(coords_list, dim=0),\n        'edge_index': torch.cat(edge_index_list, dim=1) if any(e.numel() > 0 for e in edge_index_list) else torch.zeros(2, 0, dtype=torch.long),\n        'edge_type': torch.cat(edge_type_list, dim=0) if any(e.numel() > 0 for e in edge_type_list) else torch.zeros(0, dtype=torch.long),\n        'batch': torch.cat(batch_idx_list, dim=0),\n        'lengths': torch.tensor(lengths, dtype=torch.long),\n        'sequences': sequences,\n        'ptr': torch.tensor([0] + list(np.cumsum(lengths)), dtype=torch.long)\n    }\n\n\n# =============================================================================\n# E(3) EQUIVARIANT GRAPH NEURAL NETWORK LAYERS\n# =============================================================================\n\nclass E3InvariantEdgeFeatures(nn.Module):\n    \"\"\"Compute E(3) invariant edge features from coordinates.\"\"\"\n    \n    def __init__(self, num_rbf: int = 16, cutoff: float = 20.0):\n        super().__init__()\n        self.num_rbf = num_rbf\n        self.cutoff = cutoff\n        \n        # Register buffers on CPU first, they'll be moved with the model\n        centers = torch.linspace(0, cutoff, num_rbf)\n        widths = torch.full((num_rbf,), cutoff / num_rbf)\n        self.register_buffer('centers', centers)\n        self.register_buffer('widths', widths)\n    \n    def forward(self, coords: torch.Tensor, edge_index: torch.Tensor) -> torch.Tensor:\n        if edge_index.numel() == 0:\n            return torch.zeros(0, self.num_rbf, device=coords.device, dtype=coords.dtype)\n        \n        src, dst = edge_index[0], edge_index[1]\n        \n        diff = coords[dst] - coords[src]\n        dist = torch.norm(diff, dim=-1, keepdim=True)\n        \n        rbf = torch.exp(-((dist - self.centers) / self.widths) ** 2)\n        \n        return rbf\n\n\nclass EquivariantGraphConv(nn.Module):\n    \"\"\"E(3) Equivariant Graph Convolution Layer.\"\"\"\n    \n    def __init__(self, hidden_dim: int, edge_dim: int, num_edge_types: int = 8,\n                 coord_update: bool = True, attention: bool = True):\n        super().__init__()\n        \n        self.hidden_dim = hidden_dim\n        self.edge_dim = edge_dim\n        self.coord_update = coord_update\n        self.attention = attention\n        \n        self.edge_type_embed = nn.Embedding(num_edge_types, edge_dim)\n        \n        self.message_mlp = nn.Sequential(\n            nn.Linear(hidden_dim * 2 + edge_dim + edge_dim, hidden_dim * 2),\n            nn.SiLU(),\n            nn.Linear(hidden_dim * 2, hidden_dim),\n            nn.SiLU()\n        )\n        \n        if attention:\n            self.attention_mlp = nn.Sequential(\n                nn.Linear(hidden_dim, 1),\n                nn.Sigmoid()\n            )\n        \n        if coord_update:\n            self.coord_mlp = nn.Sequential(\n                nn.Linear(hidden_dim, hidden_dim),\n                nn.SiLU(),\n                nn.Linear(hidden_dim, 1)\n            )\n        \n        self.node_update = nn.Sequential(\n            nn.Linear(hidden_dim * 2, hidden_dim),\n            nn.SiLU(),\n            nn.Linear(hidden_dim, hidden_dim)\n        )\n        \n        self.norm = nn.LayerNorm(hidden_dim)\n    \n    def forward(self, h: torch.Tensor, coords: torch.Tensor, edge_index: torch.Tensor,\n                edge_attr: torch.Tensor, edge_type: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:\n        num_nodes = h.shape[0]\n        \n        # Handle empty edges\n        if edge_index.numel() == 0:\n            return h, coords\n        \n        src, dst = edge_index[0], edge_index[1]\n        \n        # Clamp indices to valid range\n        src = src.clamp(0, num_nodes - 1)\n        dst = dst.clamp(0, num_nodes - 1)\n        \n        edge_type_clamped = edge_type.clamp(0, self.edge_type_embed.num_embeddings - 1)\n        edge_type_emb = self.edge_type_embed(edge_type_clamped)\n        \n        rel_pos = coords[dst] - coords[src]\n        \n        h_src = h[src]\n        h_dst = h[dst]\n        \n        message_input = torch.cat([h_src, h_dst, edge_attr, edge_type_emb], dim=-1)\n        messages = self.message_mlp(message_input)\n        \n        if self.attention:\n            attn = self.attention_mlp(messages)\n            messages = messages * attn\n        \n        h_agg = scatter_add(messages, dst, dim=0, dim_size=num_nodes)\n        \n        h_out = h + self.node_update(torch.cat([h, h_agg], dim=-1))\n        h_out = self.norm(h_out)\n        \n        coords_out = coords\n        if self.coord_update:\n            coord_weights = self.coord_mlp(messages)\n            coord_messages = rel_pos * coord_weights\n            \n            coord_agg = scatter_mean(coord_messages, dst, dim=0, dim_size=num_nodes)\n            coords_out = coords + coord_agg\n        \n        return h_out, coords_out\n\n\n# =============================================================================\n# GRAPH DIFFUSION DENOISING NETWORK\n# =============================================================================\n\nclass TimeEmbedding(nn.Module):\n    \"\"\"Timestep embedding with sinusoidal encoding.\"\"\"\n    \n    def __init__(self, dim: int):\n        super().__init__()\n        self.dim = dim\n        \n        self.mlp = nn.Sequential(\n            nn.Linear(dim, dim * 4),\n            nn.GELU(),\n            nn.Linear(dim * 4, dim)\n        )\n    \n    def forward(self, t: torch.Tensor) -> torch.Tensor:\n        half_dim = self.dim // 2\n        embeddings = math.log(10000) / (half_dim - 1)\n        embeddings = torch.exp(torch.arange(half_dim, device=t.device, dtype=torch.float32) * -embeddings)\n        embeddings = t[:, None].float() * embeddings[None, :]\n        embeddings = torch.cat([torch.sin(embeddings), torch.cos(embeddings)], dim=-1)\n        \n        return self.mlp(embeddings)\n\n\nclass GraphDiffusionNetwork(nn.Module):\n    \"\"\"Graph-based denoising network for diffusion.\"\"\"\n    \n    def __init__(self, node_input_dim: int = 37, hidden_dim: int = 256, edge_dim: int = 32,\n                 num_layers: int = 8, num_edge_types: int = 8, dropout: float = 0.1):\n        super().__init__()\n        \n        self.hidden_dim = hidden_dim\n        self.edge_dim = edge_dim\n        \n        self.node_encoder = nn.Sequential(\n            nn.Linear(node_input_dim, hidden_dim),\n            nn.LayerNorm(hidden_dim),\n            nn.SiLU(),\n            nn.Linear(hidden_dim, hidden_dim)\n        )\n        \n        self.coord_encoder = nn.Sequential(\n            nn.Linear(3, hidden_dim // 2),\n            nn.SiLU(),\n            nn.Linear(hidden_dim // 2, hidden_dim)\n        )\n        \n        self.feature_combine = nn.Linear(hidden_dim * 2, hidden_dim)\n        \n        self.edge_encoder = E3InvariantEdgeFeatures(num_rbf=edge_dim)\n        \n        self.time_embed = TimeEmbedding(hidden_dim)\n        \n        self.layers = nn.ModuleList()\n        for i in range(num_layers):\n            self.layers.append(\n                EquivariantGraphConv(\n                    hidden_dim=hidden_dim,\n                    edge_dim=edge_dim,\n                    num_edge_types=num_edge_types,\n                    coord_update=True,\n                    attention=True\n                )\n            )\n        \n        self.noise_predictor = nn.Sequential(\n            nn.Linear(hidden_dim, hidden_dim),\n            nn.SiLU(),\n            nn.Dropout(dropout),\n            nn.Linear(hidden_dim, hidden_dim // 2),\n            nn.SiLU(),\n            nn.Linear(hidden_dim // 2, 3)\n        )\n        \n        self.norm = nn.LayerNorm(hidden_dim)\n    \n    def forward(self, noisy_coords: torch.Tensor, timesteps: torch.Tensor,\n                node_features: torch.Tensor, edge_index: torch.Tensor,\n                edge_type: torch.Tensor, batch: torch.Tensor) -> torch.Tensor:\n        \n        h = self.node_encoder(node_features)\n        coord_features = self.coord_encoder(noisy_coords)\n        \n        h = self.feature_combine(torch.cat([h, coord_features], dim=-1))\n        \n        t_emb = self.time_embed(timesteps)\n        h = h + t_emb[batch]\n        \n        # Handle empty edges\n        if edge_index.numel() == 0:\n            h = self.norm(h)\n            return self.noise_predictor(h)\n        \n        edge_attr = self.edge_encoder(noisy_coords, edge_index)\n        \n        coords = noisy_coords\n        for layer in self.layers:\n            h, coords = layer(h, coords, edge_index, edge_attr, edge_type)\n            edge_attr = self.edge_encoder(coords, edge_index)\n        \n        h = self.norm(h)\n        noise_pred = self.noise_predictor(h)\n        \n        return noise_pred\n\n\n# =============================================================================\n# GRAPH DIFFUSION PROCESS\n# =============================================================================\n\nclass GraphDiffusion:\n    \"\"\"Diffusion process adapted for graph-structured data.\"\"\"\n    \n    def __init__(self, num_timesteps: int = 1000, beta_start: float = 1e-4,\n                 beta_end: float = 0.02, schedule: str = 'cosine',\n                 device: torch.device = None):\n        \n        self.num_timesteps = num_timesteps\n        self.device = device if device is not None else torch.device('cpu')\n        \n        # Create schedule on CPU first\n        if schedule == 'linear':\n            betas = torch.linspace(beta_start, beta_end, num_timesteps, dtype=torch.float32)\n        elif schedule == 'cosine':\n            betas = self._cosine_schedule(num_timesteps)\n        else:\n            raise ValueError(f\"Unknown schedule: {schedule}\")\n        \n        # Compute all derived tensors on CPU\n        self.betas = betas\n        self.alphas = 1.0 - betas\n        self.alphas_cumprod = torch.cumprod(self.alphas, dim=0)\n        self.alphas_cumprod_prev = F.pad(self.alphas_cumprod[:-1], (1, 0), value=1.0)\n        \n        self.sqrt_alphas_cumprod = torch.sqrt(self.alphas_cumprod)\n        self.sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - self.alphas_cumprod)\n        \n        self.posterior_variance = betas * (1.0 - self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod + 1e-8)\n        self.posterior_mean_coef1 = betas * torch.sqrt(self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod + 1e-8)\n        self.posterior_mean_coef2 = (1.0 - self.alphas_cumprod_prev) * torch.sqrt(self.alphas) / (1.0 - self.alphas_cumprod + 1e-8)\n    \n    def _cosine_schedule(self, num_timesteps: int, s: float = 0.008) -> torch.Tensor:\n        steps = torch.linspace(0, num_timesteps, num_timesteps + 1, dtype=torch.float32)\n        f_t = torch.cos((steps / num_timesteps + s) / (1 + s) * math.pi / 2) ** 2\n        alphas_cumprod = f_t / f_t[0]\n        betas = 1 - alphas_cumprod[1:] / alphas_cumprod[:-1]\n        return torch.clamp(betas, 0.0001, 0.999)\n    \n    def to(self, device: torch.device):\n        \"\"\"Move all tensors to device.\"\"\"\n        self.device = device\n        self.betas = self.betas.to(device)\n        self.alphas = self.alphas.to(device)\n        self.alphas_cumprod = self.alphas_cumprod.to(device)\n        self.alphas_cumprod_prev = self.alphas_cumprod_prev.to(device)\n        self.sqrt_alphas_cumprod = self.sqrt_alphas_cumprod.to(device)\n        self.sqrt_one_minus_alphas_cumprod = self.sqrt_one_minus_alphas_cumprod.to(device)\n        self.posterior_variance = self.posterior_variance.to(device)\n        self.posterior_mean_coef1 = self.posterior_mean_coef1.to(device)\n        self.posterior_mean_coef2 = self.posterior_mean_coef2.to(device)\n        return self\n    \n    def q_sample(self, x_0: torch.Tensor, t: torch.Tensor, \n                 batch: torch.Tensor, noise: torch.Tensor = None) -> torch.Tensor:\n        \"\"\"Forward diffusion q(x_t | x_0).\"\"\"\n        if noise is None:\n            noise = torch.randn_like(x_0)\n        \n        # Ensure indices are valid\n        t = t.clamp(0, self.num_timesteps - 1)\n        \n        # Get coefficients - move to same device as x_0\n        sqrt_alpha = self.sqrt_alphas_cumprod.to(x_0.device)[t][batch].unsqueeze(-1)\n        sqrt_one_minus_alpha = self.sqrt_one_minus_alphas_cumprod.to(x_0.device)[t][batch].unsqueeze(-1)\n        \n        return sqrt_alpha * x_0 + sqrt_one_minus_alpha * noise\n    \n    def p_mean_variance(self, model: nn.Module, x_t: torch.Tensor, t: torch.Tensor,\n                        node_features: torch.Tensor, edge_index: torch.Tensor,\n                        edge_type: torch.Tensor, batch: torch.Tensor) -> Dict[str, torch.Tensor]:\n        \"\"\"Compute mean and variance for reverse process.\"\"\"\n        \n        noise_pred = model(x_t, t, node_features, edge_index, edge_type, batch)\n        \n        t = t.clamp(0, self.num_timesteps - 1)\n        device = x_t.device\n        \n        sqrt_alpha = self.sqrt_alphas_cumprod.to(device)[t][batch].unsqueeze(-1)\n        sqrt_one_minus_alpha = self.sqrt_one_minus_alphas_cumprod.to(device)[t][batch].unsqueeze(-1)\n        \n        x_0_pred = (x_t - sqrt_one_minus_alpha * noise_pred) / (sqrt_alpha + 1e-8)\n        x_0_pred = torch.clamp(x_0_pred, -100, 100)\n        \n        coef1 = self.posterior_mean_coef1.to(device)[t][batch].unsqueeze(-1)\n        coef2 = self.posterior_mean_coef2.to(device)[t][batch].unsqueeze(-1)\n        \n        posterior_mean = coef1 * x_0_pred + coef2 * x_t\n        posterior_var = self.posterior_variance.to(device)[t][batch].unsqueeze(-1)\n        \n        return {\n            'mean': posterior_mean,\n            'variance': posterior_var,\n            'x_0_pred': x_0_pred,\n            'noise_pred': noise_pred\n        }\n    \n    @torch.no_grad()\n    def p_sample(self, model: nn.Module, x_t: torch.Tensor, t: torch.Tensor,\n                 node_features: torch.Tensor, edge_index: torch.Tensor,\n                 edge_type: torch.Tensor, batch: torch.Tensor) -> torch.Tensor:\n        \"\"\"Sample x_{t-1} from p(x_{t-1} | x_t).\"\"\"\n        \n        out = self.p_mean_variance(model, x_t, t, node_features, edge_index, edge_type, batch)\n        \n        noise = torch.randn_like(x_t)\n        nonzero_mask = (t[batch] > 0).float().unsqueeze(-1)\n        \n        x_prev = out['mean'] + nonzero_mask * torch.sqrt(out['variance'].clamp(min=1e-8)) * noise\n        \n        return x_prev\n    \n    @torch.no_grad()\n    def ddim_sample(self, model: nn.Module, x_t: torch.Tensor, t: torch.Tensor,\n                    t_prev: torch.Tensor, node_features: torch.Tensor,\n                    edge_index: torch.Tensor, edge_type: torch.Tensor,\n                    batch: torch.Tensor, eta: float = 0.0) -> torch.Tensor:\n        \"\"\"DDIM sampling step.\"\"\"\n        \n        noise_pred = model(x_t, t, node_features, edge_index, edge_type, batch)\n        \n        t = t.clamp(0, self.num_timesteps - 1)\n        t_prev = t_prev.clamp(0, self.num_timesteps - 1)\n        device = x_t.device\n        \n        alpha_t = self.alphas_cumprod.to(device)[t][batch].unsqueeze(-1)\n        alpha_prev = self.alphas_cumprod.to(device)[t_prev][batch].unsqueeze(-1)\n        \n        x_0_pred = (x_t - torch.sqrt(1 - alpha_t + 1e-8) * noise_pred) / (torch.sqrt(alpha_t) + 1e-8)\n        x_0_pred = torch.clamp(x_0_pred, -100, 100)\n        \n        sigma_sq = (1 - alpha_prev) / (1 - alpha_t + 1e-8) * (1 - alpha_t / (alpha_prev + 1e-8))\n        sigma = eta * torch.sqrt(torch.clamp(sigma_sq, min=0))\n        \n        dir_xt = torch.sqrt(torch.clamp(1 - alpha_prev - sigma**2, min=0)) * noise_pred\n        \n        x_prev = torch.sqrt(alpha_prev) * x_0_pred + dir_xt\n        \n        if eta > 0:\n            noise = torch.randn_like(x_t)\n            x_prev = x_prev + sigma * noise\n        \n        return x_prev\n    \n    @torch.no_grad()\n    def sample(self, model: nn.Module, node_features: torch.Tensor,\n               edge_index: torch.Tensor, edge_type: torch.Tensor,\n               batch: torch.Tensor, num_nodes: int,\n               num_steps: int = None, use_ddim: bool = True,\n               eta: float = 0.0) -> torch.Tensor:\n        \"\"\"Generate samples from noise.\"\"\"\n        \n        model.eval()\n        device = node_features.device\n        B = int(batch.max().item()) + 1\n        \n        x_t = torch.randn(num_nodes, 3, device=device)\n        \n        if use_ddim and num_steps is not None:\n            timesteps = torch.linspace(self.num_timesteps - 1, 0, num_steps).long()\n        else:\n            timesteps = torch.arange(self.num_timesteps - 1, -1, -1)\n        \n        for i, t in enumerate(timesteps):\n            t_batch = torch.full((B,), t.item(), device=device, dtype=torch.long)\n            \n            if use_ddim and i < len(timesteps) - 1:\n                t_prev = timesteps[i + 1]\n                t_prev_batch = torch.full((B,), t_prev.item(), device=device, dtype=torch.long)\n                x_t = self.ddim_sample(model, x_t, t_batch, t_prev_batch,\n                                       node_features, edge_index, edge_type, batch, eta)\n            else:\n                x_t = self.p_sample(model, x_t, t_batch, node_features,\n                                    edge_index, edge_type, batch)\n        \n        return x_t\n\n\n# =============================================================================\n# TRAINING\n# =============================================================================\n\nclass GraphDiffusionTrainer:\n    \"\"\"Training class for graph diffusion model.\"\"\"\n    \n    def __init__(self, model: nn.Module, diffusion: GraphDiffusion,\n                 train_dataset: Dataset, valid_dataset: Dataset = None,\n                 batch_size: int = 16, learning_rate: float = 1e-4,\n                 num_epochs: int = 100, gradient_accumulation: int = 1,\n                 device: torch.device = None):\n        \n        self.device = device if device is not None else DEVICE\n        self.model = model.to(self.device)\n        self.diffusion = diffusion\n        self.num_epochs = num_epochs\n        self.gradient_accumulation = gradient_accumulation\n        \n        self.train_loader = DataLoader(\n            train_dataset, batch_size=batch_size, shuffle=True,\n            collate_fn=collate_graph_batch, num_workers=0, pin_memory=False\n        )\n        \n        self.valid_loader = None\n        if valid_dataset and len(valid_dataset) > 0:\n            self.valid_loader = DataLoader(\n                valid_dataset, batch_size=batch_size, shuffle=False,\n                collate_fn=collate_graph_batch, num_workers=0, pin_memory=False\n            )\n        \n        self.optimizer = AdamW(model.parameters(), lr=learning_rate, weight_decay=0.01)\n        \n        total_steps = max(len(self.train_loader) * num_epochs // gradient_accumulation, 1)\n        self.scheduler = OneCycleLR(\n            self.optimizer,\n            max_lr=learning_rate,\n            total_steps=total_steps,\n            pct_start=0.1,\n            anneal_strategy='cos'\n        )\n        \n        self.ema_decay = 0.999\n        self.ema_model = None\n    \n    def _update_ema(self):\n        \"\"\"Update exponential moving average of model parameters.\"\"\"\n        if self.ema_model is None:\n            self.ema_model = {name: param.clone().detach()\n                            for name, param in self.model.named_parameters()}\n        else:\n            for name, param in self.model.named_parameters():\n                self.ema_model[name].mul_(self.ema_decay).add_(param.data, alpha=1 - self.ema_decay)\n    \n    def train_step(self, batch: Dict[str, torch.Tensor]) -> float:\n        \"\"\"Single training step.\"\"\"\n        self.model.train()\n        \n        coords = batch['coords'].to(self.device)\n        node_features = batch['node_features'].to(self.device)\n        edge_index = batch['edge_index'].to(self.device)\n        edge_type = batch['edge_type'].to(self.device)\n        batch_idx = batch['batch'].to(self.device)\n        \n        B = int(batch_idx.max().item()) + 1\n        \n        t = torch.randint(0, self.diffusion.num_timesteps, (B,), device=self.device)\n        \n        noise = torch.randn_like(coords)\n        \n        noisy_coords = self.diffusion.q_sample(coords, t, batch_idx, noise)\n        \n        noise_pred = self.model(noisy_coords, t, node_features, edge_index, edge_type, batch_idx)\n        \n        loss = F.mse_loss(noise_pred, noise)\n        \n        loss = loss / self.gradient_accumulation\n        loss.backward()\n        \n        return loss.item() * self.gradient_accumulation\n    \n    @torch.no_grad()\n    def validate(self) -> float:\n        \"\"\"Validation step.\"\"\"\n        if self.valid_loader is None:\n            return 0.0\n        \n        self.model.eval()\n        total_loss = 0\n        count = 0\n        \n        for batch in self.valid_loader:\n            coords = batch['coords'].to(self.device)\n            node_features = batch['node_features'].to(self.device)\n            edge_index = batch['edge_index'].to(self.device)\n            edge_type = batch['edge_type'].to(self.device)\n            batch_idx = batch['batch'].to(self.device)\n            \n            B = int(batch_idx.max().item()) + 1\n            t = torch.randint(0, self.diffusion.num_timesteps, (B,), device=self.device)\n            \n            noise = torch.randn_like(coords)\n            noisy_coords = self.diffusion.q_sample(coords, t, batch_idx, noise)\n            noise_pred = self.model(noisy_coords, t, node_features, edge_index, edge_type, batch_idx)\n            \n            loss = F.mse_loss(noise_pred, noise)\n            total_loss += loss.item() * B\n            count += B\n        \n        return total_loss / count if count > 0 else 0.0\n    \n    def train(self) -> Dict[str, List[float]]:\n        \"\"\"Full training loop.\"\"\"\n        history = {'train_loss': [], 'valid_loss': []}\n        \n        best_valid_loss = float('inf')\n        \n        for epoch in range(self.num_epochs):\n            epoch_loss = 0\n            count = 0\n            \n            self.optimizer.zero_grad()\n            \n            pbar = tqdm(self.train_loader, desc=f'Epoch {epoch+1}/{self.num_epochs}')\n            for step, batch in enumerate(pbar):\n                try:\n                    loss = self.train_step(batch)\n                    epoch_loss += loss\n                    count += 1\n                    \n                    if (step + 1) % self.gradient_accumulation == 0:\n                        torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1.0)\n                        self.optimizer.step()\n                        self.scheduler.step()\n                        self.optimizer.zero_grad()\n                        self._update_ema()\n                    \n                    pbar.set_postfix({'loss': f'{loss:.4f}', 'lr': f'{self.scheduler.get_last_lr()[0]:.2e}'})\n                except RuntimeError as e:\n                    print(f\"Error in batch: {e}\")\n                    self.optimizer.zero_grad()\n                    continue\n            \n            avg_train_loss = epoch_loss / count if count > 0 else 0.0\n            avg_valid_loss = self.validate()\n            \n            history['train_loss'].append(avg_train_loss)\n            history['valid_loss'].append(avg_valid_loss)\n            \n            if avg_valid_loss < best_valid_loss and avg_valid_loss > 0:\n                best_valid_loss = avg_valid_loss\n            \n            print(f'Epoch {epoch+1}: train_loss={avg_train_loss:.4f}, valid_loss={avg_valid_loss:.4f}')\n        \n        return history\n\n\n# =============================================================================\n# STRUCTURE REFINEMENT\n# =============================================================================\n\nclass GraphStructureRefiner:\n    \"\"\"Refines generated structures using physical constraints.\"\"\"\n    \n    def __init__(self, backbone_length: float = 5.9, min_distance: float = 3.0,\n                 base_pair_distance: float = 6.5, iterations: int = 20):\n        self.backbone_length = backbone_length\n        self.min_distance = min_distance\n        self.base_pair_distance = base_pair_distance\n        self.iterations = iterations\n    \n    def refine(self, coords: np.ndarray, sequence: str) -> np.ndarray:\n        \"\"\"Apply physical constraints to refine structure.\"\"\"\n        coords = coords.copy().astype(np.float64)\n        n = len(coords)\n        \n        for _ in range(self.iterations):\n            for i in range(n - 1):\n                vec = coords[i + 1] - coords[i]\n                dist = np.linalg.norm(vec) + 1e-8\n                if abs(dist - self.backbone_length) > 0.1:\n                    correction = (dist - self.backbone_length) / dist * 0.3\n                    coords[i] += vec * correction\n                    coords[i + 1] -= vec * correction\n            \n            for i in range(n):\n                for j in range(i + 2, n):\n                    vec = coords[j] - coords[i]\n                    dist = np.linalg.norm(vec) + 1e-8\n                    if dist < self.min_distance:\n                        correction = (self.min_distance - dist) / dist * 0.2\n                        coords[i] -= vec * correction\n                        coords[j] += vec * correction\n        \n        coords = coords - coords.mean(axis=0)\n        \n        return coords.astype(np.float32)\n\n\n# =============================================================================\n# MAIN PREDICTOR\n# =============================================================================\n\nclass RNAGraphDiffusionPredictor:\n    \"\"\"Main class for RNA structure prediction using graph diffusion.\"\"\"\n    \n    def __init__(self, model: nn.Module = None, diffusion: GraphDiffusion = None,\n                 device: torch.device = None, hidden_dim: int = 256, num_layers: int = 8):\n        \n        self.device = device if device is not None else DEVICE\n        self.graph_builder = RNAGraphBuilder()\n        self.refiner = GraphStructureRefiner()\n        \n        if model is None:\n            model = GraphDiffusionNetwork(\n                node_input_dim=37,\n                hidden_dim=hidden_dim,\n                edge_dim=32,\n                num_layers=num_layers,\n                num_edge_types=8,\n                dropout=0.1\n            )\n        \n        if diffusion is None:\n            # Create diffusion on CPU first, then move to device\n            diffusion = GraphDiffusion(\n                num_timesteps=1000,\n                schedule='cosine',\n                device=torch.device('cpu')\n            )\n        \n        self.model = model.to(self.device)\n        self.diffusion = diffusion\n    \n    def train(self, train_seqs: pd.DataFrame, train_labels: pd.DataFrame,\n              valid_seqs: pd.DataFrame = None, valid_labels: pd.DataFrame = None,\n              num_epochs: int = 50, batch_size: int = 16, learning_rate: float = 1e-4):\n        \"\"\"Train the model.\"\"\"\n        \n        train_coords = build_coords_dict(train_labels)\n        train_dataset = RNAGraphDataset(train_seqs, train_coords)\n        \n        valid_dataset = None\n        if valid_seqs is not None and valid_labels is not None:\n            valid_coords = build_coords_dict(valid_labels)\n            valid_dataset = RNAGraphDataset(valid_seqs, valid_coords)\n        \n        trainer = GraphDiffusionTrainer(\n            model=self.model,\n            diffusion=self.diffusion,\n            train_dataset=train_dataset,\n            valid_dataset=valid_dataset,\n            batch_size=batch_size,\n            learning_rate=learning_rate,\n            num_epochs=num_epochs,\n            device=self.device\n        )\n        \n        history = trainer.train()\n        \n        return history\n    \n    @torch.no_grad()\n    def predict(self, sequence: str, num_samples: int = 5,\n                num_steps: int = 100, eta: float = 0.0,\n                refine: bool = True) -> List[np.ndarray]:\n        \"\"\"Generate structure predictions.\"\"\"\n        \n        self.model.eval()\n        \n        predictions = []\n        \n        for i in range(num_samples):\n            torch.manual_seed(i * 1000 + int(time.time()) % 1000)\n            \n            graph = self.graph_builder.build_graph(sequence, coords=None)\n            \n            node_features = graph['node_features'].to(self.device)\n            edge_index = graph['edge_index'].to(self.device)\n            edge_type = graph['edge_type'].to(self.device)\n            \n            n = graph['num_nodes']\n            batch = torch.zeros(n, dtype=torch.long, device=self.device)\n            \n            sample_eta = eta + 0.05 * i\n            \n            coords = self.diffusion.sample(\n                self.model,\n                node_features,\n                edge_index,\n                edge_type,\n                batch,\n                num_nodes=n,\n                num_steps=num_steps,\n                use_ddim=True,\n                eta=sample_eta\n            )\n            \n            coords_np = coords.cpu().numpy()\n            \n            if refine:\n                coords_np = self.refiner.refine(coords_np, sequence)\n            \n            predictions.append(coords_np)\n        \n        return predictions\n    \n    def save(self, path: str):\n        \"\"\"Save model checkpoint.\"\"\"\n        torch.save({\n            'model_state_dict': self.model.state_dict(),\n        }, path)\n    \n    def load(self, path: str):\n        \"\"\"Load model checkpoint.\"\"\"\n        checkpoint = torch.load(path, map_location=self.device, weights_only=True)\n        self.model.load_state_dict(checkpoint['model_state_dict'])\n\n\n# =============================================================================\n# SUBMISSION PIPELINE\n# =============================================================================\n\ndef predict_rna_structures(predictor: RNAGraphDiffusionPredictor,\n                          sequence: str, target_id: str,\n                          n_predictions: int = 5) -> List[np.ndarray]:\n    \"\"\"Predict structures using the trained model.\"\"\"\n    \n    predictions = predictor.predict(\n        sequence,\n        num_samples=n_predictions,\n        num_steps=100,\n        eta=0.0,\n        refine=True\n    )\n    \n    return predictions\n\n\ndef run_prediction_pipeline(test_seqs: pd.DataFrame,\n                           predictor: RNAGraphDiffusionPredictor,\n                           output_file: str = 'submission.csv') -> pd.DataFrame:\n    \"\"\"Run prediction pipeline for submission.\"\"\"\n    \n    all_predictions = []\n    start_time = time.time()\n    total_targets = len(test_seqs)\n    \n    for idx, row in tqdm(test_seqs.iterrows(), total=total_targets, desc='Predicting'):\n        target_id = row['target_id']\n        sequence = row['sequence']\n        \n        try:\n            predictions = predict_rna_structures(predictor, sequence, target_id, n_predictions=5)\n        except Exception as e:\n            print(f\"Error predicting {target_id}: {e}\")\n            # Generate random predictions as fallback\n            n = len(sequence)\n            predictions = [np.random.randn(n, 3).astype(np.float32) * 10 for _ in range(5)]\n        \n        for j in range(len(sequence)):\n            pred_row = {\n                'ID': f\"{target_id}_{j + 1}\",\n                'resname': sequence[j],\n                'resid': j + 1\n            }\n            \n            for i in range(5):\n                pred_row[f'x_{i + 1}'] = float(predictions[i][j][0])\n                pred_row[f'y_{i + 1}'] = float(predictions[i][j][1])\n                pred_row[f'z_{i + 1}'] = float(predictions[i][j][2])\n            \n            all_predictions.append(pred_row)\n    \n    submission_df = pd.DataFrame(all_predictions)\n    \n    column_order = ['ID', 'resname', 'resid']\n    for i in range(1, 6):\n        for coord in ['x', 'y', 'z']:\n            column_order.append(f'{coord}_{i}')\n    submission_df = submission_df[column_order]\n    \n    submission_df.to_csv(output_file, index=False)\n    \n    print(f\"\\nGenerated predictions for {total_targets} sequences\")\n    print(f\"Total time: {time.time() - start_time:.1f}s\")\n    print(f\"Saved to: {output_file}\")\n    \n    return submission_df\n\n\n# =============================================================================\n# MAIN\n# =============================================================================\n\nif __name__ == \"__main__\":\n    DATA_DIR = '/kaggle/input/stanford-rna-3d-folding-2'\n    \n    HIDDEN_DIM = 192\n    NUM_LAYERS = 6\n    NUM_EPOCHS = 15\n    BATCH_SIZE = 16\n    LEARNING_RATE = 5e-5\n    \n    print(\"=\"*60)\n    print(\"RNA 3D Structure Prediction with Graph Diffusion Model\")\n    print(\"=\"*60)\n    \n    train_seqs, valid_seqs, test_seqs, train_labels, valid_labels = load_competition_data(DATA_DIR)\n    \n    # Create predictor with explicit CPU diffusion initialization\n    print(\"\\nInitializing model...\")\n    predictor = RNAGraphDiffusionPredictor(\n        hidden_dim=HIDDEN_DIM,\n        num_layers=NUM_LAYERS,\n        device=DEVICE\n    )\n    \n    print(f\"Model parameters: {sum(p.numel() for p in predictor.model.parameters()):,}\")\n    \n    print(\"\\nTraining graph diffusion model...\")\n    history = predictor.train(\n        train_seqs=train_seqs,\n        train_labels=train_labels,\n        valid_seqs=valid_seqs,\n        valid_labels=valid_labels,\n        num_epochs=NUM_EPOCHS,\n        batch_size=BATCH_SIZE,\n        learning_rate=LEARNING_RATE\n    )\n    \n    predictor.save('rna_graph_diffusion_model.pt')\n    print(\"Model saved to rna_graph_diffusion_model.pt\")\n    \n    print(\"\\nGenerating test predictions...\")\n    submission = run_prediction_pipeline(\n        test_seqs=test_seqs,\n        predictor=predictor,\n        output_file='submission.csv'\n    )\n    \n    print(\"\\nSubmission preview:\")\n    print(submission.head(10))\n    \n    print(\"\\nDone!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-11T23:11:11.489256Z","iopub.execute_input":"2026-01-11T23:11:11.489592Z","iopub.status.idle":"2026-01-11T23:24:30.484465Z","shell.execute_reply.started":"2026-01-11T23:11:11.489563Z","shell.execute_reply":"2026-01-11T23:24:30.483216Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}