{"cells":[{"cell_type":"markdown","metadata":{},"source":"# RNA 3D Baseline Inference - 5 Atoms\n\nLocal CV: 0.2462"},{"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 tqdm import tqdm"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"class RNABaselineModel(nn.Module):\n    def __init__(self, vocab_size=4, hidden_dim=256, num_layers=3, dropout=0.1, num_atoms=5):\n        super().__init__()\n        self.embedding = nn.Embedding(vocab_size, hidden_dim)\n        layers = []\n        for i in range(num_layers):\n            layers.extend([nn.Linear(hidden_dim, hidden_dim), nn.LayerNorm(hidden_dim), nn.ReLU(), nn.Dropout(dropout)])\n        self.mlp = nn.Sequential(*layers)\n        self.output = nn.Linear(hidden_dim, num_atoms * 3)\n        self.num_atoms = num_atoms\n    def forward(self, x):\n        batch_size, seq_len = x.shape\n        out = self.output(self.mlp(self.embedding(x)))\n        return out.view(batch_size, seq_len, self.num_atoms, 3)\n\ndef sequence_to_ids(seq):\n    return np.array([{'A':0,'C':1,'G':2,'U':3}.get(n,0) for n in seq])"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'\nmodel = RNABaselineModel(num_atoms=5)\nmodel.load_state_dict(torch.load('/kaggle/input/datasets/sunljdata/rna-3d-baseline-model/baseline_5atoms_v2.pth', map_location=device))\nmodel.to(device).eval()\nprint('Model loaded')"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"test_df = pd.read_csv('/kaggle/input/competitions/stanford-rna-3d-folding-2/test_sequences.csv')\nprint(f'Test: {len(test_df)} structures')"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"rows = []\nfor _, r in tqdm(test_df.iterrows(), total=len(test_df)):\n    seq = r['sequence']\n    with torch.no_grad():\n        coords = model(torch.LongTensor(sequence_to_ids(seq)).unsqueeze(0).to(device)).squeeze(0).cpu().numpy()\n    for i, nuc in enumerate(seq):\n        row = {'ID': f\"{r['target_id']}_{i+1}\", 'resname': nuc, 'resid': i+1}\n        for a in range(5):\n            row.update({f'x_{a+1}': coords[i,a,0], f'y_{a+1}': coords[i,a,1], f'z_{a+1}': coords[i,a,2]})\n        rows.append(row)\nprint(f'{len(rows)} rows')"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"pd.DataFrame(rows).to_csv('submission.csv', index=False)\nprint('Done')"}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"}},"nbformat":4,"nbformat_minor":4}