{"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"}],"dockerImageVersionId":31260,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ============================================================\n# Stanford RNA 3D Folding Part 2\n# Baseline Transformer Model\n#\n# Licensed under the Apache License, Version 2.0\n# http://www.apache.org/licenses/LICENSE-2.0\n# ============================================================","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-20T06:18:27.189904Z","iopub.execute_input":"2026-01-20T06:18:27.190124Z","iopub.status.idle":"2026-01-20T06:18:27.194531Z","shell.execute_reply.started":"2026-01-20T06:18:27.190104Z","shell.execute_reply":"2026-01-20T06:18:27.193792Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport random\nimport numpy as np\nimport pandas as pd\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T07:53:03.743297Z","iopub.execute_input":"2026-01-20T07:53:03.743989Z","iopub.status.idle":"2026-01-20T07:53:03.747888Z","shell.execute_reply.started":"2026-01-20T07:53:03.74396Z","shell.execute_reply":"2026-01-20T07:53:03.74722Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def seed_everything(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nseed_everything()\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\ndevice","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T07:53:07.91764Z","iopub.execute_input":"2026-01-20T07:53:07.917967Z","iopub.status.idle":"2026-01-20T07:53:07.925972Z","shell.execute_reply.started":"2026-01-20T07:53:07.91794Z","shell.execute_reply":"2026-01-20T07:53:07.925318Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_PATH = \"/kaggle/input/stanford-rna-3d-folding-2\"\n\nprint(os.listdir(DATA_PATH))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T07:53:10.690949Z","iopub.execute_input":"2026-01-20T07:53:10.691259Z","iopub.status.idle":"2026-01-20T07:53:10.71646Z","shell.execute_reply.started":"2026-01-20T07:53:10.691232Z","shell.execute_reply":"2026-01-20T07:53:10.715886Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\ntrain_seq = pd.read_csv(f\"{DATA_PATH}/train_sequences.csv\")\ntrain_labels = pd.read_csv(f\"{DATA_PATH}/train_labels.csv\")\n\nval_seq = pd.read_csv(f\"{DATA_PATH}/validation_sequences.csv\")\nval_labels = pd.read_csv(f\"{DATA_PATH}/validation_labels.csv\")\n\ntest_seq = pd.read_csv(f\"{DATA_PATH}/test_sequences.csv\")\n\ntrain_seq.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T07:53:13.428831Z","iopub.execute_input":"2026-01-20T07:53:13.429143Z","iopub.status.idle":"2026-01-20T07:53:22.072055Z","shell.execute_reply.started":"2026-01-20T07:53:13.429116Z","shell.execute_reply":"2026-01-20T07:53:22.071375Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_labels.columns","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T07:53:25.088508Z","iopub.execute_input":"2026-01-20T07:53:25.089261Z","iopub.status.idle":"2026-01-20T07:53:25.094285Z","shell.execute_reply.started":"2026-01-20T07:53:25.089232Z","shell.execute_reply":"2026-01-20T07:53:25.093517Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Keep only required columns\ntrain_labels = train_labels[[\"ID\", \"resid\", \"x_1\", \"y_1\", \"z_1\"]]\n\n# Rename for consistency\ntrain_labels = train_labels.rename(columns={\n    \"ID\": \"target_id\",\n    \"x_1\": \"x\",\n    \"y_1\": \"y\",\n    \"z_1\": \"z\"\n})\n\ntrain_labels.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T07:53:27.550209Z","iopub.execute_input":"2026-01-20T07:53:27.550796Z","iopub.status.idle":"2026-01-20T07:53:28.113853Z","shell.execute_reply.started":"2026-01-20T07:53:27.550765Z","shell.execute_reply":"2026-01-20T07:53:28.113229Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_seq = train_seq[[\"target_id\", \"sequence\"]]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T07:53:31.281666Z","iopub.execute_input":"2026-01-20T07:53:31.282424Z","iopub.status.idle":"2026-01-20T07:53:31.286979Z","shell.execute_reply.started":"2026-01-20T07:53:31.282393Z","shell.execute_reply":"2026-01-20T07:53:31.286321Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from collections import defaultdict\n\n# 🔹 ID clean: 157D_1 → 157D\ntrain_labels[\"target_id\"] = train_labels[\"target_id\"].str.split(\"_\").str[0]\n\ncoords_dict = defaultdict(list)\n\nfor _, row in train_labels.iterrows():\n    coords_dict[row[\"target_id\"]].append(\n        [row[\"x\"], row[\"y\"], row[\"z\"]]\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T07:53:33.520959Z","iopub.execute_input":"2026-01-20T07:53:33.521278Z","iopub.status.idle":"2026-01-20T07:59:03.406116Z","shell.execute_reply.started":"2026-01-20T07:53:33.521248Z","shell.execute_reply":"2026-01-20T07:59:03.405443Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tid = train_seq.iloc[0][\"target_id\"]\n\nprint(\"Target ID:\", tid)\nprint(\"Sequence length:\", len(train_seq.iloc[0][\"sequence\"]))\nprint(\"Coords length:\", len(coords_dict[tid]))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T07:59:19.539176Z","iopub.execute_input":"2026-01-20T07:59:19.539488Z","iopub.status.idle":"2026-01-20T07:59:19.544635Z","shell.execute_reply.started":"2026-01-20T07:59:19.53946Z","shell.execute_reply":"2026-01-20T07:59:19.543808Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"WINDOW_SIZE = 512   # 🔥 KEY PARAMETER\n\ndef collate_fn(batch):\n    seqs, coords = zip(*batch)\n\n    batch_seqs = []\n    batch_coords = []\n    batch_masks = []\n\n    for seq, coord in zip(seqs, coords):\n        L = len(seq)\n\n        # Random window start (training-time augmentation)\n        if L > WINDOW_SIZE:\n            start = random.randint(0, L - WINDOW_SIZE)\n            end = start + WINDOW_SIZE\n        else:\n            start = 0\n            end = L\n\n        seq_win = seq[start:end]\n        coord_win = coord[start:end]\n\n        valid = ~torch.isnan(coord_win).any(dim=1)\n\n        batch_seqs.append(seq_win)\n        batch_coords.append(torch.nan_to_num(coord_win, nan=0.0))\n        batch_masks.append(valid)\n\n    # Padding (now max_len <= WINDOW_SIZE)\n    max_len = max(len(s) for s in batch_seqs)\n    B = len(batch_seqs)\n\n    padded_seqs = torch.zeros(B, max_len, dtype=torch.long)\n    padded_coords = torch.zeros(B, max_len, 3)\n    mask = torch.zeros(B, max_len, dtype=torch.bool)\n\n    for i in range(B):\n        l = len(batch_seqs[i])\n        padded_seqs[i, :l] = batch_seqs[i]\n        padded_coords[i, :l] = batch_coords[i]\n        mask[i, :l] = batch_masks[i]\n\n    return padded_seqs, padded_coords, mask","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T07:59:28.491896Z","iopub.execute_input":"2026-01-20T07:59:28.492195Z","iopub.status.idle":"2026-01-20T07:59:28.500383Z","shell.execute_reply.started":"2026-01-20T07:59:28.492168Z","shell.execute_reply":"2026-01-20T07:59:28.499554Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RNADataset(Dataset):\n    def __init__(self, seq_df, coords_dict):\n        self.seq_df = seq_df.reset_index(drop=True)\n        self.coords_dict = coords_dict\n\n        self.vocab = {\"A\": 0, \"U\": 1, \"G\": 2, \"C\": 3}\n\n    def encode(self, seq):\n        return torch.tensor([self.vocab[x] for x in seq], dtype=torch.long)\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        tid = row[\"target_id\"]\n\n        seq = self.encode(row[\"sequence\"])\n        coords = torch.tensor(self.coords_dict[tid], dtype=torch.float32)\n\n        return seq, coords","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T07:59:32.679513Z","iopub.execute_input":"2026-01-20T07:59:32.680134Z","iopub.status.idle":"2026-01-20T07:59:32.686311Z","shell.execute_reply.started":"2026-01-20T07:59:32.680105Z","shell.execute_reply":"2026-01-20T07:59:32.685378Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = RNADataset(train_seq, coords_dict)\n\n# Validation labels same logic apply karo\nval_labels = val_labels[[\"ID\", \"resid\", \"x_1\", \"y_1\", \"z_1\"]]\nval_labels = val_labels.rename(columns={\n    \"ID\": \"target_id\",\n    \"x_1\": \"x\",\n    \"y_1\": \"y\",\n    \"z_1\": \"z\"\n})\n\nval_labels[\"target_id\"] = val_labels[\"target_id\"].str.split(\"_\").str[0]\n\nfrom collections import defaultdict\nval_coords_dict = defaultdict(list)\n\nfor _, row in val_labels.iterrows():\n    val_coords_dict[row[\"target_id\"]].append([row[\"x\"], row[\"y\"], row[\"z\"]])\n\nval_seq = val_seq[[\"target_id\", \"sequence\"]]\nval_dataset = RNADataset(val_seq, val_coords_dict)\n\nlen(train_dataset), len(val_dataset)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T07:59:38.161275Z","iopub.execute_input":"2026-01-20T07:59:38.161943Z","iopub.status.idle":"2026-01-20T07:59:38.87742Z","shell.execute_reply.started":"2026-01-20T07:59:38.161908Z","shell.execute_reply":"2026-01-20T07:59:38.876871Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_loader = DataLoader(\n    train_dataset,\n    batch_size=4,        # SAFE\n    shuffle=True,\n    collate_fn=collate_fn,\n    num_workers=2,\n    pin_memory=True\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=4,\n    shuffle=False,\n    collate_fn=collate_fn,\n    num_workers=2,\n    pin_memory=True\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T07:59:43.124029Z","iopub.execute_input":"2026-01-20T07:59:43.124347Z","iopub.status.idle":"2026-01-20T07:59:44.06722Z","shell.execute_reply.started":"2026-01-20T07:59:43.124319Z","shell.execute_reply":"2026-01-20T07:59:44.066512Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"batch = next(iter(train_loader))\n\nseqs, coords, mask = batch\n\nprint(\"Seqs shape:\", seqs.shape)\nprint(\"Coords shape:\", coords.shape)\nprint(\"Mask shape:\", mask.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T08:00:16.230466Z","iopub.execute_input":"2026-01-20T08:00:16.231101Z","iopub.status.idle":"2026-01-20T08:00:16.501587Z","shell.execute_reply.started":"2026-01-20T08:00:16.23107Z","shell.execute_reply":"2026-01-20T08:00:16.50072Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RNATransformer(nn.Module):\n    def __init__(\n        self,\n        vocab_size=4,\n        d_model=128,\n        nhead=8,\n        num_layers=6,\n        dim_feedforward=512,\n        dropout=0.1\n    ):\n        super().__init__()\n\n        self.embedding = nn.Embedding(vocab_size, d_model)\n\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=d_model,\n            nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout,\n            batch_first=True\n        )\n\n        self.encoder = nn.TransformerEncoder(\n            encoder_layer,\n            num_layers=num_layers\n        )\n\n        self.coord_head = nn.Linear(d_model, 3)\n\n    def forward(self, seqs, mask):\n        \"\"\"\n        seqs: (B, L)\n        mask: (B, L)  True = valid\n        \"\"\"\n        x = self.embedding(seqs)            # (B, L, d_model)\n        x = self.encoder(\n            x,\n            src_key_padding_mask=~mask\n        )\n        coords = self.coord_head(x)         # (B, L, 3)\n        return coords","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T07:20:40.108346Z","iopub.execute_input":"2026-01-20T07:20:40.109208Z","iopub.status.idle":"2026-01-20T07:20:40.115639Z","shell.execute_reply.started":"2026-01-20T07:20:40.109171Z","shell.execute_reply":"2026-01-20T07:20:40.114814Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def masked_mse_loss(pred, target, mask):\n    \"\"\"\n    pred, target: (B, L, 3)\n    mask: (B, L)\n    \"\"\"\n    mask = mask.unsqueeze(-1)              # (B, L, 1)\n    loss = ((pred - target) ** 2) * mask\n    return loss.sum() / mask.sum()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T08:00:33.468465Z","iopub.execute_input":"2026-01-20T08:00:33.469275Z","iopub.status.idle":"2026-01-20T08:00:33.473948Z","shell.execute_reply.started":"2026-01-20T08:00:33.469237Z","shell.execute_reply":"2026-01-20T08:00:33.473145Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = RNATransformer().to(device)\n\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=1e-4,\n    weight_decay=1e-4\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T08:00:35.986778Z","iopub.execute_input":"2026-01-20T08:00:35.987356Z","iopub.status.idle":"2026-01-20T08:00:36.013101Z","shell.execute_reply.started":"2026-01-20T08:00:35.987328Z","shell.execute_reply":"2026-01-20T08:00:36.012597Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"seqs, coords, mask = next(iter(train_loader))\nprint(seqs.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T08:01:04.085071Z","iopub.execute_input":"2026-01-20T08:01:04.085813Z","iopub.status.idle":"2026-01-20T08:01:04.333174Z","shell.execute_reply.started":"2026-01-20T08:01:04.085781Z","shell.execute_reply":"2026-01-20T08:01:04.332288Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.train()\n\nseqs, coords, mask = next(iter(train_loader))\n\nseqs = seqs.to(device)\ncoords = coords.to(device)\nmask = mask.to(device)\n\npred = model(seqs, mask)\nloss = masked_mse_loss(pred, coords, mask)\n\nloss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T08:01:12.086808Z","iopub.execute_input":"2026-01-20T08:01:12.087646Z","iopub.status.idle":"2026-01-20T08:01:12.364912Z","shell.execute_reply.started":"2026-01-20T08:01:12.087562Z","shell.execute_reply":"2026-01-20T08:01:12.364084Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EPOCHS = 2   # enough for baseline\n\nfor epoch in range(EPOCHS):\n    torch.cuda.empty_cache()\n\n    model.train()\n    train_loss = 0.0\n\n    for seqs, coords, mask in train_loader:\n        seqs = seqs.to(device)\n        coords = coords.to(device)\n        mask = mask.to(device)\n\n        optimizer.zero_grad(set_to_none=True)\n        pred = model(seqs, mask)\n        loss = masked_mse_loss(pred, coords, mask)\n        loss.backward()\n        optimizer.step()\n\n        train_loss += loss.item()\n\n    train_loss /= len(train_loader)\n    print(f\"Epoch {epoch+1} | Train Loss: {train_loss:.2f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T08:01:59.912466Z","iopub.execute_input":"2026-01-20T08:01:59.913081Z","iopub.status.idle":"2026-01-20T08:03:16.322968Z","shell.execute_reply.started":"2026-01-20T08:01:59.913045Z","shell.execute_reply":"2026-01-20T08:03:16.322181Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval()\n\nRNA_VOCAB = {\"A\": 0, \"U\": 1, \"G\": 2, \"C\": 3}\nsubmission_rows = []\n\ntotal = len(test_seq)\n\nwith torch.no_grad():\n    for idx, row in test_seq.iterrows():\n        seq_encoded = torch.tensor(\n            [RNA_VOCAB[x] for x in row[\"sequence\"]],\n            dtype=torch.long\n        ).unsqueeze(0).to(device)\n\n        mask = torch.ones_like(seq_encoded).bool()\n\n        preds = []\n        for _ in range(5):\n            coords = model(seq_encoded, mask)\n            coords = coords.squeeze(0).cpu().numpy()\n            preds.append(coords)\n\n        for i, base in enumerate(row[\"sequence\"]):\n            out = {\n                \"ID\": row[\"target_id\"],\n                \"resname\": base,\n                \"resid\": i + 1\n            }\n            for k in range(5):\n                out[f\"x_{k+1}\"] = float(preds[k][i][0])\n                out[f\"y_{k+1}\"] = float(preds[k][i][1])\n                out[f\"z_{k+1}\"] = float(preds[k][i][2])\n\n            submission_rows.append(out)\n\n        # 🔥 progress print every 1 RNA\n        print(f\"Done {idx+1}/{total}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T08:07:13.83179Z","iopub.execute_input":"2026-01-20T08:07:13.832714Z","iopub.status.idle":"2026-01-20T08:07:14.9523Z","shell.execute_reply.started":"2026-01-20T08:07:13.832681Z","shell.execute_reply":"2026-01-20T08:07:14.95171Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission = pd.DataFrame(submission_rows)\nsubmission.to_csv(\"submission.csv\", index=False)\n\nsubmission.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T08:08:51.436301Z","iopub.execute_input":"2026-01-20T08:08:51.436805Z","iopub.status.idle":"2026-01-20T08:08:51.717402Z","shell.execute_reply.started":"2026-01-20T08:08:51.436774Z","shell.execute_reply":"2026-01-20T08:08:51.716712Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}