{"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,"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-20T10:10:34.957177Z","iopub.execute_input":"2026-01-20T10:10:34.957688Z","iopub.status.idle":"2026-01-20T10:10:34.96152Z","shell.execute_reply.started":"2026-01-20T10:10:34.957659Z","shell.execute_reply":"2026-01-20T10:10:34.960945Z"}},"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-20T10:10:37.378308Z","iopub.execute_input":"2026-01-20T10:10:37.378834Z","iopub.status.idle":"2026-01-20T10:10:41.649533Z","shell.execute_reply.started":"2026-01-20T10:10:37.378807Z","shell.execute_reply":"2026-01-20T10:10:41.648614Z"}},"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-20T10:10:44.229699Z","iopub.execute_input":"2026-01-20T10:10:44.230491Z","iopub.status.idle":"2026-01-20T10:10:44.345527Z","shell.execute_reply.started":"2026-01-20T10:10:44.230456Z","shell.execute_reply":"2026-01-20T10:10:44.344849Z"}},"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-20T10:10:47.719699Z","iopub.execute_input":"2026-01-20T10:10:47.720475Z","iopub.status.idle":"2026-01-20T10:10:47.726389Z","shell.execute_reply.started":"2026-01-20T10:10:47.720433Z","shell.execute_reply":"2026-01-20T10:10:47.725628Z"}},"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-20T10:10:50.28867Z","iopub.execute_input":"2026-01-20T10:10:50.288989Z","iopub.status.idle":"2026-01-20T10:11:01.427123Z","shell.execute_reply.started":"2026-01-20T10:10:50.288961Z","shell.execute_reply":"2026-01-20T10:11:01.426437Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_labels.columns","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T10:11:06.508195Z","iopub.execute_input":"2026-01-20T10:11:06.508618Z","iopub.status.idle":"2026-01-20T10:11:06.51427Z","shell.execute_reply.started":"2026-01-20T10:11:06.50859Z","shell.execute_reply":"2026-01-20T10:11:06.513628Z"}},"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-20T10:11:09.217066Z","iopub.execute_input":"2026-01-20T10:11:09.217356Z","iopub.status.idle":"2026-01-20T10:11:09.745186Z","shell.execute_reply.started":"2026-01-20T10:11:09.21733Z","shell.execute_reply":"2026-01-20T10:11:09.744257Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_seq = train_seq[[\"target_id\", \"sequence\"]]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T10:11:12.721988Z","iopub.execute_input":"2026-01-20T10:11:12.72269Z","iopub.status.idle":"2026-01-20T10:11:12.727341Z","shell.execute_reply.started":"2026-01-20T10:11:12.722659Z","shell.execute_reply":"2026-01-20T10:11:12.726649Z"}},"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-20T10:11:15.480579Z","iopub.execute_input":"2026-01-20T10:11:15.481258Z","iopub.status.idle":"2026-01-20T10:16:22.92665Z","shell.execute_reply.started":"2026-01-20T10:11:15.481229Z","shell.execute_reply":"2026-01-20T10:16:22.926061Z"}},"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-20T10:17:41.372169Z","iopub.execute_input":"2026-01-20T10:17:41.372719Z","iopub.status.idle":"2026-01-20T10:17:41.378033Z","shell.execute_reply.started":"2026-01-20T10:17:41.372693Z","shell.execute_reply":"2026-01-20T10:17:41.377194Z"}},"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-20T10:17:47.595065Z","iopub.execute_input":"2026-01-20T10:17:47.595608Z","iopub.status.idle":"2026-01-20T10:17:47.602847Z","shell.execute_reply.started":"2026-01-20T10:17:47.59558Z","shell.execute_reply":"2026-01-20T10:17:47.602109Z"}},"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-20T10:17:50.5439Z","iopub.execute_input":"2026-01-20T10:17:50.544723Z","iopub.status.idle":"2026-01-20T10:17:50.550005Z","shell.execute_reply.started":"2026-01-20T10:17:50.544692Z","shell.execute_reply":"2026-01-20T10:17:50.549369Z"}},"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-20T10:17:57.577339Z","iopub.execute_input":"2026-01-20T10:17:57.577989Z","iopub.status.idle":"2026-01-20T10:17:58.280767Z","shell.execute_reply.started":"2026-01-20T10:17:57.577962Z","shell.execute_reply":"2026-01-20T10:17:58.280146Z"}},"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-20T10:18:01.545012Z","iopub.execute_input":"2026-01-20T10:18:01.545353Z","iopub.status.idle":"2026-01-20T10:18:01.550608Z","shell.execute_reply.started":"2026-01-20T10:18:01.545327Z","shell.execute_reply":"2026-01-20T10:18:01.549438Z"}},"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-20T10:18:04.062402Z","iopub.execute_input":"2026-01-20T10:18:04.062739Z","iopub.status.idle":"2026-01-20T10:18:04.607786Z","shell.execute_reply.started":"2026-01-20T10:18:04.062713Z","shell.execute_reply":"2026-01-20T10:18:04.607004Z"}},"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-20T10:18:08.13432Z","iopub.execute_input":"2026-01-20T10:18:08.134611Z","iopub.status.idle":"2026-01-20T10:18:08.141328Z","shell.execute_reply.started":"2026-01-20T10:18:08.134583Z","shell.execute_reply":"2026-01-20T10:18:08.140443Z"}},"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-20T10:18:11.83408Z","iopub.execute_input":"2026-01-20T10:18:11.83458Z","iopub.status.idle":"2026-01-20T10:18:11.838693Z","shell.execute_reply.started":"2026-01-20T10:18:11.834551Z","shell.execute_reply":"2026-01-20T10:18:11.837926Z"}},"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-20T10:18:16.00176Z","iopub.execute_input":"2026-01-20T10:18:16.002363Z","iopub.status.idle":"2026-01-20T10:18:19.131229Z","shell.execute_reply.started":"2026-01-20T10:18:16.002334Z","shell.execute_reply":"2026-01-20T10:18:19.130669Z"}},"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-20T10:18:21.882255Z","iopub.execute_input":"2026-01-20T10:18:21.88304Z","iopub.status.idle":"2026-01-20T10:18:22.061776Z","shell.execute_reply.started":"2026-01-20T10:18:21.883006Z","shell.execute_reply":"2026-01-20T10:18:22.060944Z"}},"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-20T10:18:24.743468Z","iopub.execute_input":"2026-01-20T10:18:24.744125Z","iopub.status.idle":"2026-01-20T10:18:25.890357Z","shell.execute_reply.started":"2026-01-20T10:18:24.74409Z","shell.execute_reply":"2026-01-20T10:18:25.889591Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EPOCHS = 1   # 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-20T10:18:57.278655Z","iopub.execute_input":"2026-01-20T10:18:57.279398Z","iopub.status.idle":"2026-01-20T10:19:33.336271Z","shell.execute_reply.started":"2026-01-20T10:18:57.279367Z","shell.execute_reply":"2026-01-20T10:19:33.335552Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval()\n\nRNA_VOCAB = {\"A\": 0, \"U\": 1, \"G\": 2, \"C\": 3}\nsubmission_rows = []\n\nWINDOW = 512\n\nwith torch.no_grad():\n    for idx, row in test_seq.iterrows():\n        seq = row[\"sequence\"]\n        L = len(seq)\n\n        preds_all = [[] for _ in range(5)]\n\n        for start in range(0, L, WINDOW):\n            end = min(start + WINDOW, L)\n\n            seq_chunk = torch.tensor(\n                [RNA_VOCAB[x] for x in seq[start:end]],\n                dtype=torch.long\n            ).unsqueeze(0).to(device)\n\n            mask = torch.ones_like(seq_chunk).bool()\n\n            for k in range(5):\n                coords = model(seq_chunk, mask)\n                coords = coords.squeeze(0).cpu().numpy()\n                preds_all[k].append(coords)\n\n        preds_all = [np.vstack(p) for p in preds_all]\n\n        for i, base in enumerate(seq):\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_all[k][i][0])\n                out[f\"y_{k+1}\"] = float(preds_all[k][i][1])\n                out[f\"z_{k+1}\"] = float(preds_all[k][i][2])\n\n            submission_rows.append(out)\n\n        print(f\"Done {idx+1}/{len(test_seq)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T10:19:54.0629Z","iopub.execute_input":"2026-01-20T10:19:54.063598Z","iopub.status.idle":"2026-01-20T10:19:55.160764Z","shell.execute_reply.started":"2026-01-20T10:19:54.06356Z","shell.execute_reply":"2026-01-20T10:19:55.160137Z"}},"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-20T10:20:00.167075Z","iopub.execute_input":"2026-01-20T10:20:00.167743Z","iopub.status.idle":"2026-01-20T10:20:00.427442Z","shell.execute_reply.started":"2026-01-20T10:20:00.167711Z","shell.execute_reply":"2026-01-20T10:20:00.42636Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}