{"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":"none","dataSources":[{"sourceId":118765,"databundleVersionId":15231210,"sourceType":"competition"}],"dockerImageVersionId":31234,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-09T11:43:26.075115Z","iopub.execute_input":"2026-01-09T11:43:26.075713Z","iopub.status.idle":"2026-01-09T11:43:43.272883Z","shell.execute_reply.started":"2026-01-09T11:43:26.075679Z","shell.execute_reply":"2026-01-09T11:43:43.271537Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"_kg_hide-output":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Project Overview & Dataset Documentation\n\n## 1. Objective\n\nThe goal of this project is to predict the 3D structure of RNA molecules based on their nucleotide sequence. This is a fundamental problem in biology, as structure determines function (e.g., how a ribozyme catalyzes a reaction or how a standard mRNA is translated).\n\nWe will approach this as a **sequence-to-coordinate regression problem** using deep learning.\n\n---\n\n## 2. Core Sequence Files\n\nThe primary input data is located in:\n\n- `train_sequences.csv`\n- `validation_sequences.csv`\n- `test_sequences.csv`\n\nThese files contain the RNA sequences that need to be folded.\n\n### File Schema\n\n| Column Name       | Description                                                                                                                                                                                |\n| ----------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------ |\n| `target_id`       | Unique identifier for the target. In training data, this corresponds to `pdb_id_chain_id`.                                                                                                 |\n| `sequence`        | The RNA sequence to predict. This is a concatenation of all chains specified in the stoichiometry.                                                                                         |\n| `stoichiometry`   | Specifies which chains from the experiment are part of the target (e.g., `{chain:number}`).                                                                                                |\n| `all_sequences`   | FASTA-formatted sequences of all molecules in the experiment, including those not in the target (e.g., viral proteins or DNA partners). Use `extra/parse_fasta_py.py` to parse this field. |\n| `ligand_ids`      | Three-letter codes for any small molecules present. These are not predicted but may influence RNA folding.                                                                                 |\n| `ligand_SMILES`   | Chemical structures of the ligands, represented as SMILES strings.                                                                                                                         |\n| `temporal_cutoff` | Publication date of the structure, useful for preventing data leakage.                                                                                                                     |\n\n---\n\n## 3. Labels (Ground Truth)\n\nGround truth RNA structures are provided in:\n\n- `train_labels.csv`\n- `validation_labels.csv`\n\n### Structural Coordinates\n\n- The label files contain **x, y, z coordinates** for the **C1' atom** (RNA backbone) of each residue.\n\n### Multiple Conformations\n\n- Some RNA targets have multiple experimentally observed conformations.\n- In such cases, the training labels may include:\n  - `x_1, y_1, z_1`\n  - `x_2, y_2, z_2`\n  - ...\n  - `x_n, y_n, z_n`\n\n### Residue Indexing\n\n- Residue numbering (`resid`) follows **1-based indexing**.\n","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport glob\nimport os\n\n# --- Configuration ---\nBASE_DIR = '/kaggle/input/stanford-rna-3d-folding-2/'\ncsv_files = {\n    'train_seq': f'{BASE_DIR}train_sequences.csv',\n    'train_labels': f'{BASE_DIR}train_labels.csv',\n    'val_seq': f'{BASE_DIR}validation_sequences.csv',\n    'val_labels': f'{BASE_DIR}validation_labels.csv',\n    'test_seq': f'{BASE_DIR}test_sequences.csv',\n    'sample_sub': f'{BASE_DIR}sample_submission.csv'\n}\n\n# --- 1. Load and Inspect CSV Data ---\ndataframes = {}\nprint(\"--- Loading Dataframes ---\")\nfor name, path in csv_files.items():\n    if os.path.exists(path):\n        df = pd.read_csv(path)\n        dataframes[name] = df\n        print(f\"\\n{name} shape: {df.shape}\")\n        print(f\"Columns: {list(df.columns)}\")\n        print(\"-\" * 30)\n    else:\n        print(f\"Warning: {path} not found.\")\n\n# --- 2. Sequence Length Distribution ---\n# Merging sequences to compare lengths\nif 'train_seq' in dataframes and 'test_seq' in dataframes:\n    train_df = dataframes['train_seq'].copy()\n    test_df = dataframes['test_seq'].copy()\n    val_df = dataframes['val_seq'].copy()\n\n    train_df['set'] = 'Train'\n    test_df['set'] = 'Test'\n    val_df['set'] = 'Validation'\n    \n    train_df['seq_len'] = train_df['sequence'].apply(len)\n    test_df['seq_len'] = test_df['sequence'].apply(len)\n    val_df['seq_len'] = val_df['sequence'].apply(len)\n\n    all_seqs = pd.concat([train_df, val_df, test_df])\n\n    plt.figure(figsize=(12, 6))\n    sns.histplot(data=all_seqs, x='seq_len', hue='set', kde=True, bins=50)\n    plt.title(\"Distribution of RNA Sequence Lengths\")\n    plt.xlabel(\"Sequence Length (nucleotides)\")\n    plt.savefig('sequence_length_distribution.png')\n    plt.show()\n    \n    print(\"\\nSequence Length Statistics:\")\n    print(all_seqs.groupby('set')['seq_len'].describe())\n\n# --- 3. Inspect MSA Depth ---\n# Checking a few MSA files to see how many homologous sequences they contain\nmsa_files = glob.glob(f'{BASE_DIR}MSA/*.fasta')\nprint(f\"\\nFound {len(msa_files)} MSA files.\")\n\nif len(msa_files) > 0:\n    print(\"\\n--- Inspecting Sample MSA Depth ---\")\n    msa_depths = []\n    # Check first 100 files to save time\n    for fpath in msa_files[:100]:\n        with open(fpath, 'r') as f:\n            # Count lines starting with '>'\n            count = sum(1 for line in f if line.startswith('>'))\n            msa_depths.append(count)\n    \n    plt.figure(figsize=(10, 5))\n    sns.histplot(msa_depths, bins=30, kde=True)\n    plt.title(\"Distribution of MSA Depth (First 100 files)\")\n    plt.xlabel(\"Number of Homologous Sequences\")\n    plt.savefig('msa_depth_distribution.png')\n    plt.show()\n    print(f\"Average MSA Depth (Sample): {np.mean(msa_depths):.2f}\")\n\n# --- 4. Inspect Labels (Coordinates) ---\nif 'train_labels' in dataframes:\n    print(\"\\n--- Train Labels Preview ---\")\n    lbl = dataframes['train_labels']\n    # Check for residues with missing coordinates (if any logic implies it)\n    # The file has x_1, y_1, z_1 etc. Let's see basic stats of coordinates\n    coords = lbl[['x_1', 'y_1', 'z_1']].dropna()\n    print(coords.describe())\n    \n    # Plot a random structure (first one found)\n    sample_id = lbl['ID'].iloc[0].split('_')[0] # Get target ID\n    sample_struct = lbl[lbl['ID'].str.startswith(sample_id)]\n    \n    fig = plt.figure(figsize=(8, 8))\n    ax = fig.add_subplot(111, projection='3d')\n    ax.plot(sample_struct['x_1'], sample_struct['y_1'], sample_struct['z_1'], marker='o', linestyle='-', markersize=4)\n    ax.set_title(f\"3D Trace of C1' Backbone: {sample_id}\")\n    plt.savefig('sample_rna_structure.png')\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-09T11:48:08.673635Z","iopub.execute_input":"2026-01-09T11:48:08.674161Z","iopub.status.idle":"2026-01-09T11:49:16.981003Z","shell.execute_reply.started":"2026-01-09T11:48:08.674138Z","shell.execute_reply":"2026-01-09T11:49:16.979993Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# EDA Findings and Training Pipeline Roadmap\n\nBased on the Exploratory Data Analysis (EDA), the following are the **critical insights** and **immediate next steps** required to build a robust training pipeline.\n\n---\n\n## Key Insights from EDA\n\n### 1. Sequence Length Outliers\n\n- The training dataset contains **extremely large RNA sequences**, with lengths reaching up to **125,580 nucleotides**.\n- Such sequence lengths are **not feasible for standard deep learning models** due to GPU memory (OOM) constraints.\n\n**Action Items:**\n- Filter or **crop** training sequences.\n- A practical cutoff for RNA 3D structure modeling on standard GPUs is typically in the range of:\n  - **512 – 2048 tokens**\n\n---\n\n### 2. Absolute Coordinate Values\n\n- The structural coordinates have a **mean value around ~160**, indicating they are expressed in the **global PDB coordinate frame**.\n\n**Why this is a problem:**\n- Training a neural network to directly predict absolute coordinates leads to poor generalization.\n\n**Action Items:**\n- **Center the coordinates** for each structure:\n  - Subtract the mean of the x, y, z coordinates per structure.\n  - Ensure every structure’s centroid is at **(0, 0, 0)**.\n\n---\n\n### 3. Rich MSA (Multiple Sequence Alignment) Data\n\n- On average, each target has **~8,800 homologous sequences**.\n- This is a **highly valuable signal** for structure prediction.\n\n**Why this matters:**\n- Evolutionary information from MSAs is one of the **strongest predictors of RNA structure**.\n\n**Action Items:**\n- Incorporate **MSA-based features** into the model architecture.\n- Avoid sequence-only modeling where possible.\n\n---\n\n## Next Step: Data Pipeline (PyTorch)\n\nThe next logical step is to implement a **custom PyTorch `Dataset` class**.\n\n### The Dataset Class Should:\n\n1. **Tokenize RNA Sequences**\n   - Encode nucleotides as integers:\n     - `A → 0`\n     - `C → 1`\n     - `G → 2`\n     - `U → 3`\n\n2. **Load and Center 3D Coordinates**\n   - Read x, y, z coordinates for each residue.\n   - Center each structure so the centroid is at `(0, 0, 0)`.\n\n3. **Handle Variable-Length Sequences**\n   - Apply **padding** to allow batching of sequences with different lengths.\n   - Ensure padding masks are correctly handled during training.\n\n---\n\n## Outcome\n\nCompleting this data pipeline will provide:\n- Memory-efficient training samples\n- Numerically stable coordinate targets\n- A foundation for integrating **MSA-aware models** (e.g., AlphaFold-style architectures)\n\nThis sets the stage for **model architecture selection and training**.\n","metadata":{}},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset, DataLoader\nimport pandas as pd\nimport numpy as np\nimport os\n\n# --- 1. Load Data explicitly ---\nprint(\"Loading data...\")\nbase_path = '/kaggle/input/stanford-rna-3d-folding-2'\ntrain_seq = pd.read_csv(f'{base_path}/train_sequences.csv')\ntrain_labels = pd.read_csv(f'{base_path}/train_labels.csv')\n\n# Quick check to ensure alignment columns exist\nprint(f\"Loaded {len(train_seq)} sequences and {len(train_labels)} label rows.\")\n\n# --- 2. Define the Dataset Class ---\nNUC_MAP = {'A': 0, 'C': 1, 'G': 2, 'U': 3}\nPAD_TOKEN = 4\n\nclass RNADataset(Dataset):\n    def __init__(self, sequences_df, labels_df, max_len=256, split='train'):\n        self.split = split\n        self.max_len = max_len\n        \n        # Filter: Keep only sequences shorter than max_len for this test\n        # (This avoids OOM errors during debugging)\n        self.seq_df = sequences_df[sequences_df['sequence'].str.len() <= max_len].reset_index(drop=True)\n        self.labels_df = labels_df\n\n    def __len__(self):\n        return len(self.seq_df)\n\n    def __getitem__(self, idx):\n        # A. Prepare Input Sequence\n        row = self.seq_df.iloc[idx]\n        target_id = row['target_id']\n        sequence = row['sequence']\n        \n        # Tokenize (Char -> Int)\n        tokenized = [NUC_MAP.get(n, 4) for n in sequence] \n        seq_len = len(tokenized)\n        \n        # Pad inputs\n        pad_len = self.max_len - seq_len\n        tokenized = tokenized + [PAD_TOKEN] * pad_len\n        mask = [1] * seq_len + [0] * pad_len \n        \n        input_tensor = torch.tensor(tokenized, dtype=torch.long)\n        mask_tensor = torch.tensor(mask, dtype=torch.bool)\n        \n        # B. Prepare Output Coordinates (Labels)\n        if self.split == 'train':\n            # Match target_id to the labels file\n            # ID format in labels is usually: {target_id}_{residue_index}\n            # We filter for rows where ID starts with the target_id\n            subset = self.labels_df[self.labels_df['ID'].str.startswith(f\"{target_id}_\")]\n            \n            # If subset is empty, handle gracefully (return dummy or skip)\n            if len(subset) == 0:\n                # Fallback for debugging if ID match fails\n                return input_tensor, mask_tensor, torch.zeros((self.max_len, 3))\n            \n            # Extract coordinates\n            coords = subset[['x_1', 'y_1', 'z_1']].values\n            \n            # Ensure coords length matches sequence length (crop if necessary)\n            coords = coords[:self.max_len]\n            \n            # CENTER COORDINATES (Important!)\n            if len(coords) > 0:\n                coords = coords - np.mean(coords, axis=0)\n            \n            # Pad coordinates\n            pad_coords = np.zeros((self.max_len, 3))\n            pad_coords[:len(coords), :] = coords\n            \n            label_tensor = torch.tensor(pad_coords, dtype=torch.float32)\n            \n            return input_tensor, mask_tensor, label_tensor\n            \n        return input_tensor, mask_tensor\n\n# --- 3. Initialize and Test ---\nprint(\"Initializing Dataset...\")\n# We use a small max_len (256) for a quick test run\ntrain_ds = RNADataset(train_seq, train_labels, max_len=256, split='train')\ntrain_loader = DataLoader(train_ds, batch_size=8, shuffle=True)\n\n# Fetch one batch to verify\ninputs, masks, labels = next(iter(train_loader))\n\nprint(\"-\" * 30)\nprint(\"SUCCESS!\")\nprint(f\"Batch Input Shape: {inputs.shape}  (Batch Size, Max Len)\")\nprint(f\"Batch Label Shape: {labels.shape} (Batch Size, Max Len, 3)\")\nprint(f\"Sample Input (First 10 tokens): {inputs[0][:10]}\")\nprint(\"-\" * 30)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-09T11:57:25.852537Z","iopub.execute_input":"2026-01-09T11:57:25.852975Z","iopub.status.idle":"2026-01-09T11:57:51.339858Z","shell.execute_reply.started":"2026-01-09T11:57:25.852938Z","shell.execute_reply":"2026-01-09T11:57:51.33883Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 1: The MSA Helper Function\n\nFirst, we need a robust function to parse the .fasta files into a numerical tensor.\n\n**Constraint**: Some MSAs have 8000+ sequences. We cannot feed all of them into a GPU. We must randomly sample a subset (e.g., 32 or 64) during training.\n","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport torch\n\n# Configuration\nMAX_MSA_DEPTH = 32  # Number of sequences to sample from the MSA\nMAX_SEQ_LEN = 256   # Cropping length\nNUC_MAP = {'A': 0, 'C': 1, 'G': 2, 'U': 3, '-': 4} # '-' is a gap in MSA\n\ndef parse_msa(msa_path, max_depth=MAX_MSA_DEPTH, max_len=MAX_SEQ_LEN):\n    \"\"\"\n    Reads an MSA file and converts it to a tensor of shape (Depth, Len).\n    \"\"\"\n    if not os.path.exists(msa_path):\n        # Return dummy MSA if file missing\n        return np.full((max_depth, max_len), 4, dtype=np.int32)\n\n    sequences = []\n    current_seq = []\n    \n    with open(msa_path, 'r') as f:\n        for line in f:\n            line = line.strip()\n            if line.startswith(\">\"):\n                if current_seq:\n                    sequences.append(\"\".join(current_seq))\n                    current_seq = []\n            else:\n                current_seq.append(line)\n        if current_seq:\n            sequences.append(\"\".join(current_seq))\n\n    # Convert to Integers\n    msa_tensor = []\n    for seq in sequences:\n        # Tokenize (and crop if too long)\n        indices = [NUC_MAP.get(c, 4) for c in seq[:max_len]]\n        # Pad if too short\n        if len(indices) < max_len:\n            indices += [4] * (max_len - len(indices))\n        msa_tensor.append(indices)\n\n    # Convert to Numpy\n    msa_matrix = np.array(msa_tensor, dtype=np.int32)\n\n    # Downsample if too deep (Random sampling is best for training)\n    current_depth = msa_matrix.shape[0]\n    if current_depth > max_depth:\n        indices = np.random.choice(current_depth, max_depth, replace=False)\n        msa_matrix = msa_matrix[indices, :]\n    elif current_depth < max_depth:\n        # Pad with gaps if too shallow\n        padding = np.full((max_depth - current_depth, max_len), 4, dtype=np.int32)\n        msa_matrix = np.concatenate([msa_matrix, padding], axis=0)\n\n    return msa_matrix","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-09T12:18:42.112462Z","iopub.execute_input":"2026-01-09T12:18:42.11302Z","iopub.status.idle":"2026-01-09T12:18:42.124846Z","shell.execute_reply.started":"2026-01-09T12:18:42.112991Z","shell.execute_reply":"2026-01-09T12:18:42.123355Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 2: Upgraded Dataset Class\n\nWe integrate the MSA parser into the Dataset. The model will now receive (Batch, Depth, Len) inputs.\n","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader\nimport pandas as pd\n\nclass RNADatasetMSA(Dataset):\n    def __init__(self, sequences_df, labels_df, msa_dir, max_len=256, split='train'):\n        self.split = split\n        self.max_len = max_len\n        self.msa_dir = msa_dir\n        # Filter out sequences that are way too long for this specific demo\n        self.seq_df = sequences_df[sequences_df['sequence'].str.len() <= max_len].reset_index(drop=True)\n        self.labels_df = labels_df\n\n    def __len__(self):\n        return len(self.seq_df)\n\n    def __getitem__(self, idx):\n        row = self.seq_df.iloc[idx]\n        target_id = row['target_id']\n        \n        # 1. Load MSA (Keep existing logic)\n        pdb_id = target_id.split('_')[0] \n        msa_path = os.path.join(self.msa_dir, f\"{pdb_id}.MSA.fasta\")\n        msa_data = parse_msa(msa_path, max_len=self.max_len)\n        msa_tensor = torch.tensor(msa_data, dtype=torch.long)\n\n        # 2. Load Coordinates with NaN Handling\n        label_tensor = torch.zeros((self.max_len, 3))\n        # Mask: 1 if we have coords, 0 if padding OR missing coords\n        loss_mask = torch.zeros(self.max_len, dtype=torch.bool) \n        \n        if self.split == 'train':\n            subset = self.labels_df[self.labels_df['ID'].str.startswith(f\"{target_id}_\")]\n            \n            if len(subset) > 0:\n                # Extract raw values\n                raw_coords = subset[['x_1', 'y_1', 'z_1']].values[:self.max_len]\n                \n                # Check for NaNs in this specific target's coords\n                valid_mask = ~np.isnan(raw_coords).any(axis=1) # Rows without NaNs\n                \n                if np.sum(valid_mask) > 0:\n                    # Only calculate centroid using VALID coordinates\n                    valid_coords = raw_coords[valid_mask]\n                    centroid = np.mean(valid_coords, axis=0)\n                    \n                    # Center the valid coordinates\n                    centered_coords = valid_coords - centroid\n                    \n                    # Fill the tensor\n                    # We need to map back to the original indices (0 to len)\n                    # For simplicity in this demo, we assume consecutive valid rows match sequence\n                    # In production, you must match 'resid' column to sequence index\n                    \n                    # Safe fill:\n                    seq_len = min(len(raw_coords), self.max_len)\n                    \n                    # Create a temporary holder\n                    temp_coords = np.zeros((seq_len, 3))\n                    # Where raw_coords was valid, fill with centered. Where NaN, leave 0.\n                    # (Note: This is a simplification. Ideally we align by residue ID)\n                    \n                    # Easier robust approach for the demo: \n                    # Replace NaNs in raw_coords with the Centroid (so they become 0 after centering)\n                    safe_coords = np.nan_to_num(raw_coords, nan=np.nanmean(raw_coords, axis=0))\n                    \n                    # Re-calculate mean of safe coords\n                    final_centroid = np.mean(safe_coords, axis=0)\n                    final_coords = safe_coords - final_centroid\n                    \n                    label_tensor[:seq_len] = torch.tensor(final_coords, dtype=torch.float32)\n                    loss_mask[:seq_len] = True\n\n        return msa_tensor, label_tensor, loss_mask","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-09T12:23:54.537013Z","iopub.execute_input":"2026-01-09T12:23:54.538641Z","iopub.status.idle":"2026-01-09T12:23:54.55088Z","shell.execute_reply.started":"2026-01-09T12:23:54.538603Z","shell.execute_reply":"2026-01-09T12:23:54.549331Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Step 3: The Model (Triangle Attention)\n\nWe will implement a simplified version of the **Evoformer** block from AlphaFold 2. This architecture is designed specifically for folding problems.\n\n### Core Concepts\n\n1.  **MSA Input:** $ (Batch, Depth, Length) $\n\n    - We process the raw evolutionary data into an embedding.\n\n2.  **Pair Representation:** $ (Batch, Length, Length, Channels) $\n\n    - Unlike standard Transformers which work on a list of tokens, folding models work on a **Pair Map**.\n    - This map represents the spatial relationship between residue $i$ and residue $j$.\n\n3.  **Triangle Attention:**\n    - Standard attention relates $i$ to all $j$.\n    - Triangle attention updates the edge $(i, j)$ by looking at a third node $k$ to form a triangle.\n    - **Intuition:** If residue $i$ is close to $k$, and $j$ is close to $k$, then $i$ and $j$ are likely constrained relative to each other.\n\n$$ \\text{Update}(i, j) \\leftarrow \\sum_k \\text{Attention}(i, k) \\times \\text{Attention}(k, j) $$\n\n4.  **Structure Module:**\n    - Finally, we project the Pair Representation down to 3D coordinates $(x, y, z)$.\n","metadata":{}},{"cell_type":"code","source":"import torch.nn as nn\nimport torch.nn.functional as F\n\nclass TriangleAttentionBlock(nn.Module):\n    def __init__(self, d_pair, heads=4):\n        super().__init__()\n        self.d_pair = d_pair\n        self.heads = heads\n        self.scale = (d_pair // heads) ** -0.5\n        \n        # Projections for Q, K, V\n        self.to_q = nn.Linear(d_pair, d_pair)\n        self.to_k = nn.Linear(d_pair, d_pair)\n        self.to_v = nn.Linear(d_pair, d_pair)\n        \n        # Gate for \"Multiplicative\" update (like in AlphaFold)\n        self.gate = nn.Linear(d_pair, d_pair)\n        self.out_proj = nn.Linear(d_pair, d_pair)\n        self.norm = nn.LayerNorm(d_pair)\n\n    def forward(self, pair_rep):\n        \"\"\"\n        pair_rep: (Batch, Len, Len, Channels)\n        \"\"\"\n        B, L, _, C = pair_rep.shape\n        residual = pair_rep\n        \n        # 1. Norm\n        x = self.norm(pair_rep)\n        \n        # 2. Q, K, V Projections\n        q = self.to_q(x).view(B, L, L, self.heads, C // self.heads)\n        k = self.to_k(x).view(B, L, L, self.heads, C // self.heads)\n        v = self.to_v(x).view(B, L, L, self.heads, C // self.heads)\n        \n        # 3. Triangle Attention Logic (Simplified)\n        # We want to relate edge (i,j) using edges (i,k) and (k,j)\n        # Permute for matrix multiplication\n        q = q.permute(0, 3, 1, 2, 4) # (B, H, L, L, head_dim)\n        k = k.permute(0, 3, 1, 2, 4)\n        v = v.permute(0, 3, 1, 2, 4)\n        \n        # Attention scores: Q(i,k) * K(k,j) -> Map (i,j)\n        attn = torch.matmul(q, k.transpose(-2, -1)) * self.scale\n        attn = F.softmax(attn, dim=-1)\n        \n        # Aggregate: Attn(i,j) * V(k,j)\n        out = torch.matmul(attn, v) # (B, H, L, L, head_dim)\n        \n        # Reshape back\n        out = out.permute(0, 2, 3, 1, 4).reshape(B, L, L, C)\n        \n        # Gating\n        g = torch.sigmoid(self.gate(x))\n        out = self.out_proj(out) * g\n        \n        return residual + out\n\nclass RNAFoldingModel(nn.Module):\n    def __init__(self, d_msa=32, d_pair=64):\n        super().__init__()\n        \n        # 1. Embedding: MSA -> Pair Representation\n        # Simple trick: Outer product of MSA features to create (L, L) map\n        self.msa_proj = nn.Linear(5, d_pair) # 5 = vocab size (A,C,G,U,Gap)\n        self.pair_proj = nn.Linear(d_pair * 2, d_pair)\n        \n        # 2. Trunk: Triangle Attention Blocks\n        self.triangle_block = TriangleAttentionBlock(d_pair)\n        \n        # 3. Structure Module (Predict XYZ)\n        # We predict (L, 3) from the Pair Map (L, L, C)\n        # We pool over the second dimension: \"What represents residue i relative to all others?\"\n        self.to_coords = nn.Sequential(\n            nn.Linear(d_pair, 32),\n            nn.ReLU(),\n            nn.Linear(32, 3)\n        )\n\n    def forward(self, msa):\n        # msa: (Batch, Depth, Len)\n        B, D, L = msa.shape\n        \n        # Take the first sequence (query sequence) for simple embedding features\n        # One-hot encode the query sequence (Batch, Len, 5)\n        query_seq = F.one_hot(msa[:, 0, :], num_classes=5).float()\n        \n        # Create Initial Pair Representation (Outer Concatenation)\n        # (B, L, C) -> (B, L, 1, C) and (B, 1, L, C)\n        x = self.msa_proj(query_seq) # (B, L, d_pair)\n        pair_rep = torch.cat([\n            x.unsqueeze(2).expand(-1, -1, L, -1),\n            x.unsqueeze(1).expand(-1, L, -1, -1)\n        ], dim=-1) # (B, L, L, d_pair*2)\n        \n        pair_rep = self.pair_proj(pair_rep) # (B, L, L, d_pair)\n        \n        # Apply Triangle Attention\n        pair_rep = self.triangle_block(pair_rep)\n        \n        # Structure Prediction\n        # Mean pool over the columns (L, L, C) -> (L, C)\n        single_rep = pair_rep.mean(dim=2) \n        \n        # Project to 3D coordinates\n        coords = self.to_coords(single_rep) # (B, L, 3)\n        \n        return coords","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-09T12:23:57.833679Z","iopub.execute_input":"2026-01-09T12:23:57.834658Z","iopub.status.idle":"2026-01-09T12:23:57.85181Z","shell.execute_reply.started":"2026-01-09T12:23:57.834626Z","shell.execute_reply":"2026-01-09T12:23:57.85031Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Run Verification\n","metadata":{}},{"cell_type":"code","source":"# --- Setup ---\nmsa_dir = '/kaggle/input/stanford-rna-3d-folding-2/MSA/'\nbase_path = '/kaggle/input/stanford-rna-3d-folding-2'\n\n# Re-load data if needed (assuming train_seq and train_labels exist from previous step)\n# train_seq = pd.read_csv(...) \n\n# Initialize Upgraded Dataset\n# Re-initialize everything with the fixed class\nmsa_ds = RNADatasetMSA(train_seq, train_labels, msa_dir=msa_dir, max_len=64)\nmsa_loader = DataLoader(msa_ds, batch_size=2, shuffle=True)\n\nmodel = RNAFoldingModel(d_msa=32, d_pair=64)\n\n# Run Batch\nprint(\"Running Forward Pass (Attempt 2)...\")\nmsa_batch, label_batch, loss_mask = next(iter(msa_loader))\noutput = model(msa_batch)\n\n# Masked Loss Calculation\n# Only calculate MSE on valid residues (loss_mask == True)\n# We flatten the tensors to make masking easier\nactive_outputs = output[loss_mask]\nactive_labels = label_batch[loss_mask]\n\nif len(active_outputs) > 0:\n    loss = F.mse_loss(active_outputs, active_labels)\n    print(f\"Fixed MSE Loss: {loss.item()}\")\nelse:\n    print(\"Warning: No valid residues in this batch (bad batch).\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-09T12:24:26.052065Z","iopub.execute_input":"2026-01-09T12:24:26.052921Z","iopub.status.idle":"2026-01-09T12:24:29.893369Z","shell.execute_reply.started":"2026-01-09T12:24:26.052888Z","shell.execute_reply":"2026-01-09T12:24:29.892244Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 1. The Training Loop\n\nHere is a complete training script. It includes an optimizer (Adam), a progress bar, and checkpoints to save the best model. I've set MAX_LEN to 128 for this demo, but you should increase it to 256 or 512 on a GPU.\n","metadata":{}},{"cell_type":"code","source":"import torch.optim as optim\nfrom tqdm import tqdm\n\n# --- Hyperparameters ---\nBATCH_SIZE = 4\nLEARNING_RATE = 1e-3\nEPOCHS = 3\nMAX_LEN = 128  # Increase this if your GPU allows!\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# --- 1. Setup Data & Model ---\nprint(f\"Training on {DEVICE}...\")\nmsa_ds = RNADatasetMSA(train_seq, train_labels, msa_dir=msa_dir, max_len=MAX_LEN)\ntrain_loader = DataLoader(msa_ds, batch_size=BATCH_SIZE, shuffle=True)\n\nmodel = RNAFoldingModel(d_msa=32, d_pair=64).to(DEVICE)\noptimizer = optim.Adam(model.parameters(), lr=LEARNING_RATE)\n\n# --- 2. Training Loop ---\nmodel.train()\nfor epoch in range(EPOCHS):\n    total_loss = 0\n    progress_bar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{EPOCHS}\")\n    \n    for msa_batch, label_batch, loss_mask in progress_bar:\n        # Move to GPU\n        msa_batch = msa_batch.to(DEVICE)\n        label_batch = label_batch.to(DEVICE)\n        loss_mask = loss_mask.to(DEVICE)\n        \n        # Forward\n        optimizer.zero_grad()\n        output = model(msa_batch) # (B, L, 3)\n        \n        # Masked Loss\n        active_outputs = output[loss_mask]\n        active_labels = label_batch[loss_mask]\n        \n        if len(active_outputs) > 0:\n            loss = F.mse_loss(active_outputs, active_labels)\n            loss.backward()\n            optimizer.step()\n            \n            total_loss += loss.item()\n            progress_bar.set_postfix({'loss': f\"{loss.item():.2f}\"})\n        \n    avg_loss = total_loss / len(train_loader)\n    print(f\"Epoch {epoch+1} Complete. Average Loss: {avg_loss:.4f}\")\n\n# --- 3. Save Model ---\ntorch.save(model.state_dict(), \"rna_folding_model.pth\")\nprint(\"Model saved successfully!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-09T12:26:19.484346Z","iopub.execute_input":"2026-01-09T12:26:19.484916Z","iopub.status.idle":"2026-01-09T18:01:25.80165Z","shell.execute_reply.started":"2026-01-09T12:26:19.484888Z","shell.execute_reply":"2026-01-09T18:01:25.800312Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 2. The Submission Pipeline\n\nThe competition requires 5 predictions per sequence. Since our current model is deterministic (it always outputs the same answer for the same input), we will simply replicate our single prediction 5 times.\n\nTo generate diversity (which boosts score), you would later add \"Dropout\" to the model or train 5 slightly different models (Ensemble).\n","metadata":{}},{"cell_type":"code","source":"# --- Load Test Data ---\ntest_seq = pd.read_csv(f'{base_path}/test_sequences.csv')\n\n# Reuse the Dataset logic for Test (no labels needed)\nclass RNATestDataset(Dataset):\n    def __init__(self, sequences_df, msa_dir, max_len=MAX_LEN):\n        self.seq_df = sequences_df\n        self.msa_dir = msa_dir\n        self.max_len = max_len\n\n    def __len__(self):\n        return len(self.seq_df)\n\n    def __getitem__(self, idx):\n        row = self.seq_df.iloc[idx]\n        target_id = row['target_id']\n        sequence = row['sequence']\n        \n        # Parse MSA (Handle missing files gracefully)\n        # Note: In inference, if MSA is missing, we use a dummy (gaps)\n        pdb_id = target_id.split('_')[0] \n        msa_path = os.path.join(self.msa_dir, f\"{pdb_id}.MSA.fasta\")\n        msa_data = parse_msa(msa_path, max_len=self.max_len)\n        \n        return torch.tensor(msa_data, dtype=torch.long), target_id, len(sequence)\n\n# --- Inference Loop ---\ntest_ds = RNATestDataset(test_seq, msa_dir=msa_dir, max_len=MAX_LEN)\ntest_loader = DataLoader(test_ds, batch_size=1, shuffle=False)\n\nmodel.eval()\nsubmission_rows = []\n\nprint(\"Generating predictions...\")\nwith torch.no_grad():\n    for msa_batch, target_ids, original_lens in tqdm(test_loader):\n        msa_batch = msa_batch.to(DEVICE)\n        output = model(msa_batch) # (1, MaxLen, 3)\n        \n        coords = output[0].cpu().numpy()\n        target_id = target_ids[0]\n        seq_len = original_lens.item()\n        \n        # Crop back to original length\n        # (The prediction usually needs to match the requested residues exactly)\n        coords = coords[:seq_len]\n        \n        # Format for CSV: We need 5 sets of coordinates\n        # Since model is deterministic, we copy coords 5 times\n        # ID, resname, resid, x_1, y_1, z_1 ... x_5, y_5, z_5\n        \n        # We need to map residues. For simplicity, we assume 1-based indexing 1..L\n        for i in range(len(coords)):\n            resid = i + 1\n            x, y, z = coords[i]\n            \n            # Format row\n            # Note: resname is required. We need to fetch it from the sequence df ideally.\n            # Here we just use 'A' as a placeholder if strictly needed, \n            # OR we fetch the char from the test_seq dataframe.\n            # Let's verify format. The example shows: \"ID,resname,resid,x_1...\"\n            \n            # Construct row ID\n            row_id = f\"{target_id}_{resid}\"\n            \n            # 5 copies\n            row_data = [row_id, 'A', resid] # 'A' is dummy resname\n            for _ in range(5):\n                row_data.extend([f\"{x:.3f}\", f\"{y:.3f}\", f\"{z:.3f}\"])\n            \n            submission_rows.append(row_data)\n\n# --- Create DataFrame and Save ---\ncolumns = ['ID', 'resname', 'resid']\nfor i in range(1, 6):\n    columns.extend([f'x_{i}', f'y_{i}', f'z_{i}'])\n\nsub_df = pd.DataFrame(submission_rows, columns=columns)\nsub_df.to_csv('submission.csv', index=False)\nprint(\"submission.csv created!\")\nprint(sub_df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-09T18:18:44.250583Z","iopub.execute_input":"2026-01-09T18:18:44.253306Z","iopub.status.idle":"2026-01-09T18:18:46.773034Z","shell.execute_reply.started":"2026-01-09T18:18:44.253059Z","shell.execute_reply":"2026-01-09T18:18:46.771621Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}