{"cells":[{"cell_type":"markdown","metadata":{},"source":"# RNA 3D Structure Prediction - Baseline Inference\n\nThis notebook performs inference using the trained baseline MLP model.\n\n**Model**: Simple MLP (Embedding → 3×MLP → Output)\n**Local CV Score**: 0.1446 ± 0.0344"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"import torch\nimport torch.nn as nn\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nfrom tqdm import tqdm"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# Model Definition\nclass RNABaselineModel(nn.Module):\n    def __init__(self, vocab_size=4, hidden_dim=256, num_layers=3, dropout=0.1, coord_scale=100.0):\n        super().__init__()\n        self.vocab_size = vocab_size\n        self.hidden_dim = hidden_dim\n        self.coord_scale = coord_scale\n        \n        self.embedding = nn.Embedding(vocab_size, hidden_dim)\n        \n        layers = []\n        for i in range(num_layers):\n            layers.append(nn.Linear(hidden_dim, hidden_dim))\n            layers.append(nn.LayerNorm(hidden_dim))\n            layers.append(nn.ReLU())\n            layers.append(nn.Dropout(dropout))\n        \n        self.mlp = nn.Sequential(*layers)\n        self.output = nn.Linear(hidden_dim, 3)\n    \n    def forward(self, sequence_ids):\n        x = self.embedding(sequence_ids)\n        x = self.mlp(x)\n        coords = self.output(x) * self.coord_scale\n        return coords\n\ndef sequence_to_ids(sequence):\n    mapping = {'A': 0, 'C': 1, 'G': 2, 'U': 3}\n    return np.array([mapping.get(nuc, 0) for nuc in sequence])"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# Load model\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(f'Using device: {device}')\n\nmodel = RNABaselineModel()\nmodel.load_state_dict(torch.load('/kaggle/input/rna-3d-baseline-model/baseline_fold0.pth', map_location=device))\nmodel.to(device)\nmodel.eval()\nprint('✅ Model loaded successfully')"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# Load test data\ntest_sequences = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv')\nprint(f'Test structures: {len(test_sequences)}')\ntest_sequences.head()"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# Generate predictions\nsubmission_rows = []\n\nfor idx, row in tqdm(test_sequences.iterrows(), total=len(test_sequences), desc='Predicting'):\n    struct_id = row['target_id']\n    sequence = row['sequence']\n    \n    # Convert sequence to IDs\n    seq_ids = sequence_to_ids(sequence)\n    seq_tensor = torch.LongTensor(seq_ids).unsqueeze(0).to(device)\n    \n    # Predict\n    with torch.no_grad():\n        pred_coords = model(seq_tensor).squeeze(0).cpu().numpy()  # (n_residues, 3)\n    \n    # Replicate to 5 atoms (baseline: all atoms at same position)\n    n_residues = len(sequence)\n    coords = np.zeros((n_residues, 5, 3))\n    for i in range(5):\n        coords[:, i, :] = pred_coords\n    \n    # Create submission rows\n    for resid in range(n_residues):\n        row_id = f\"{struct_id}_{resid}\"\n        sub_row = {'ID': row_id}\n        \n        for atom_idx in range(5):\n            sub_row[f'x_{atom_idx+1}'] = coords[resid, atom_idx, 0]\n            sub_row[f'y_{atom_idx+1}'] = coords[resid, atom_idx, 1]\n            sub_row[f'z_{atom_idx+1}'] = coords[resid, atom_idx, 2]\n        \n        submission_rows.append(sub_row)\n\nprint(f'\\n✅ Generated {len(submission_rows):,} prediction rows')"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# Create submission file\nsubmission_df = pd.DataFrame(submission_rows)\nsubmission_df.to_csv('submission.csv', index=False)\n\nprint(f'✅ Submission saved: submission.csv')\nprint(f'   Shape: {submission_df.shape}')\nprint(f'   Columns: {list(submission_df.columns)}')\nsubmission_df.head()"}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.0"}},"nbformat":4,"nbformat_minor":4}