{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":[{"sourceType":"competition","sourceId":118765,"databundleVersionId":15231210}],"dockerImageVersionId":31287,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"\n# RNA 3D Folding – Stable Baseline + 5‑Structure Submission\n\nThis notebook includes:\n\n- Clean dataset pipeline\n- Transformer encoder model\n- Pairwise distance loss\n- Stable training loop\n- Sliding window inference (overlap averaging)\n- 5 stochastic structure predictions per target\n- Submission file generation (x_1..z_5)\n\nFully Kaggle submission-ready.\n","metadata":{}},{"cell_type":"markdown","source":"## 1. Imports & Configuration","metadata":{}},{"cell_type":"code","source":"\nimport os\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nimport pandas as pd\n\nfrom torch.utils.data import Dataset, DataLoader\n\nMAX_LEN = 1024\nBATCH_SIZE = 8\nEPOCHS = 20\nPATIENCE = 6\nLR = 3e-5\n\nNUC_MAP = {\"A\":0,\"U\":1,\"G\":2,\"C\":3}\nPAD_TOKEN = 4\nVOCAB_SIZE = 5\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(\"Device:\", DEVICE)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T08:56:06.077796Z","iopub.execute_input":"2026-02-26T08:56:06.078103Z","iopub.status.idle":"2026-02-26T08:56:10.5751Z","shell.execute_reply.started":"2026-02-26T08:56:06.078054Z","shell.execute_reply":"2026-02-26T08:56:10.574295Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Load Kaggle Data","metadata":{}},{"cell_type":"code","source":"\nROOT = \"/kaggle/input/competitions/stanford-rna-3d-folding-2\"\n\ntrain_seq = pd.read_csv(os.path.join(ROOT, \"train_sequences.csv\"))\ntrain_lab = pd.read_csv(os.path.join(ROOT, \"train_labels.csv\"), low_memory=False)\n\nval_seq = pd.read_csv(os.path.join(ROOT, \"validation_sequences.csv\"))\nval_lab = pd.read_csv(os.path.join(ROOT, \"validation_labels.csv\"), low_memory=False)\n\ntest_seq = pd.read_csv(os.path.join(ROOT, \"test_sequences.csv\"))\n\ntrain_lab[\"target_id\"] = train_lab[\"ID\"].str.split(\"_\").str[0]\nval_lab[\"target_id\"] = val_lab[\"ID\"].str.split(\"_\").str[0]\n\ntrain_lab[[\"x_1\",\"y_1\",\"z_1\"]] = train_lab[[\"x_1\",\"y_1\",\"z_1\"]].astype(np.float32)\nval_lab[[\"x_1\",\"y_1\",\"z_1\"]] = val_lab[[\"x_1\",\"y_1\",\"z_1\"]].astype(np.float32)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T08:56:10.576634Z","iopub.execute_input":"2026-02-26T08:56:10.576979Z","iopub.status.idle":"2026-02-26T08:56:49.456503Z","shell.execute_reply.started":"2026-02-26T08:56:10.576956Z","shell.execute_reply":"2026-02-26T08:56:49.455883Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Dataset Module","metadata":{}},{"cell_type":"code","source":"\nclass RNADataset(Dataset):\n\n    def __init__(self, seq_df, lab_df=None):\n        self.samples = []\n        seq_map = dict(zip(seq_df[\"target_id\"], seq_df[\"sequence\"]))\n\n        if lab_df is not None:\n            for tid, g in lab_df.groupby(\"target_id\"):\n\n                if tid not in seq_map:\n                    continue\n\n                full_seq = seq_map[tid]\n                L_total = len(full_seq)\n\n                copies = sorted(g[\"copy\"].unique())\n                n_copies = len(copies)\n\n                if L_total % n_copies != 0:\n                    continue\n\n                L_copy = L_total // n_copies\n\n                for i, c in enumerate(copies):\n\n                    g_c = g[g[\"copy\"] == c]\n\n                    coords = np.zeros((L_copy, 3), dtype=np.float32)\n                    mask = np.zeros(L_copy, dtype=np.float32)\n\n                    r = g_c[\"resid\"].astype(int).to_numpy() - 1\n                    xyz = g_c[[\"x_1\",\"y_1\",\"z_1\"]].to_numpy(np.float32)\n\n                    good = np.isfinite(xyz).all(axis=1)\n                    good &= (np.abs(xyz).max(axis=1) < 1e4)\n\n                    r = r[good]\n                    xyz = xyz[good]\n\n                    ok = (r >= 0) & (r < L_copy)\n                    coords[r[ok]] = xyz[ok]\n                    mask[r[ok]] = 1.0\n\n                    start = i * L_copy\n                    end = start + L_copy\n                    seq_slice = full_seq[start:end]\n\n                    if mask.sum() == 0:\n                        continue\n\n                    self.samples.append((seq_slice, coords, mask))\n\n        else:\n            for _, row in seq_df.iterrows():\n                self.samples.append((row[\"sequence\"], None, None))\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n\n        seq, coords, mask = self.samples[idx]\n\n        tokens = np.array([NUC_MAP[s] for s in seq], dtype=np.int64)\n        tokens = tokens[:MAX_LEN]\n        L = len(tokens)\n\n        if L < MAX_LEN:\n            tokens = np.pad(tokens, (0, MAX_LEN-L), constant_values=PAD_TOKEN)\n\n        tokens = torch.tensor(tokens, dtype=torch.long)\n\n        if coords is None:\n            return tokens\n\n        coords = coords[:MAX_LEN]\n        mask = mask[:MAX_LEN]\n\n        if L < MAX_LEN:\n            coords = np.pad(coords, ((0, MAX_LEN-L),(0,0)))\n            mask = np.pad(mask, (0, MAX_LEN-L))\n\n        valid = mask == 1\n        if valid.sum() > 0:\n            center = coords[valid].mean(axis=0, keepdims=True)\n            coords[valid] -= center\n\n        coords = coords / 50.0\n\n        return tokens, torch.tensor(mask, dtype=torch.float32), torch.tensor(coords, dtype=torch.float32)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T08:56:49.457417Z","iopub.execute_input":"2026-02-26T08:56:49.457674Z","iopub.status.idle":"2026-02-26T08:56:49.470687Z","shell.execute_reply.started":"2026-02-26T08:56:49.457652Z","shell.execute_reply":"2026-02-26T08:56:49.47013Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Transformer Model","metadata":{}},{"cell_type":"code","source":"\nclass RNATransformer(nn.Module):\n\n    def __init__(self, d_model=192, nhead=8, num_layers=4, dropout=0.2):\n        super().__init__()\n\n        self.embedding = nn.Embedding(VOCAB_SIZE, d_model)\n        self.pos_embedding = nn.Parameter(torch.zeros(1, MAX_LEN, d_model))\n\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=d_model,\n            nhead=nhead,\n            dim_feedforward=512,\n            dropout=dropout,\n            batch_first=True\n        )\n\n        self.encoder = nn.TransformerEncoder(encoder_layer, num_layers)\n        self.head = nn.Linear(d_model, 3)\n\n    def forward(self, tokens, mask):\n\n        x = self.embedding(tokens)\n        x = x + self.pos_embedding[:, :x.size(1), :]\n\n        padding_mask = (mask == 0)\n        x = self.encoder(x, src_key_padding_mask=padding_mask)\n\n        return self.head(x)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T08:56:57.738238Z","iopub.execute_input":"2026-02-26T08:56:57.738579Z","iopub.status.idle":"2026-02-26T08:56:57.744473Z","shell.execute_reply.started":"2026-02-26T08:56:57.73855Z","shell.execute_reply":"2026-02-26T08:56:57.743677Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Pairwise Distance Loss","metadata":{}},{"cell_type":"code","source":"\ndef pairwise_distances(coords, mask):\n\n    diff = coords.unsqueeze(2) - coords.unsqueeze(1)\n    dist = torch.sqrt((diff ** 2).sum(-1) + 1e-8)\n\n    valid = mask.unsqueeze(1) * mask.unsqueeze(2)\n\n    return dist, valid\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T08:57:02.287727Z","iopub.execute_input":"2026-02-26T08:57:02.288486Z","iopub.status.idle":"2026-02-26T08:57:02.292723Z","shell.execute_reply.started":"2026-02-26T08:57:02.28843Z","shell.execute_reply":"2026-02-26T08:57:02.291962Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Training","metadata":{}},{"cell_type":"code","source":"train_dataset = RNADataset(train_seq, train_lab)\nval_dataset   = RNADataset(val_seq, val_lab)\n\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True)\nval_loader   = DataLoader(val_dataset, batch_size=BATCH_SIZE)\n\nmodel = RNATransformer().to(DEVICE)\noptimizer = torch.optim.AdamW(model.parameters(), lr=LR)\n\nbest_val_loss = float(\"inf\")\npatience_counter = 0\n\nfor epoch in range(EPOCHS):\n\n    # --------------------\n    # TRAIN\n    # --------------------\n    model.train()\n    train_total = 0\n    train_batches = 0\n\n    for tokens, mask, coords in train_loader:\n\n        tokens = tokens.to(DEVICE)\n        mask   = mask.to(DEVICE)\n        coords = coords.to(DEVICE)\n\n        optimizer.zero_grad()\n\n        pred = model(tokens, mask)\n\n        pred_dist, valid = pairwise_distances(pred, mask)\n        true_dist, _     = pairwise_distances(coords, mask)\n\n        diff = F.smooth_l1_loss(pred_dist, true_dist, reduction=\"none\")\n        diff = diff * valid\n\n        denom = valid.sum()\n        if denom == 0:\n            continue\n\n        loss = diff.sum() / denom\n\n        if torch.isnan(loss) or torch.isinf(loss):\n            continue\n\n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        optimizer.step()\n\n        train_total += loss.item()\n        train_batches += 1\n\n    train_loss = train_total / max(train_batches, 1)\n\n    # --------------------\n    # VALIDATION\n    # --------------------\n    model.eval()\n    val_total = 0\n    val_batches = 0\n\n    with torch.no_grad():\n        for tokens, mask, coords in val_loader:\n\n            tokens = tokens.to(DEVICE)\n            mask   = mask.to(DEVICE)\n            coords = coords.to(DEVICE)\n\n            pred = model(tokens, mask)\n\n            pred_dist, valid = pairwise_distances(pred, mask)\n            true_dist, _     = pairwise_distances(coords, mask)\n\n            diff = F.smooth_l1_loss(pred_dist, true_dist, reduction=\"none\")\n            diff = diff * valid\n\n            loss = diff.sum() / (valid.sum() + 1e-8)\n\n            val_total += loss.item()\n            val_batches += 1\n\n    val_loss = val_total / max(val_batches, 1)\n\n    print(f\"Epoch {epoch+1} | Train: {train_loss:.4f} | Val: {val_loss:.4f}\")\n\n    # --------------------\n    # EARLY STOPPING\n    # --------------------\n    if val_loss < best_val_loss:\n        best_val_loss = val_loss\n        patience_counter = 0\n        torch.save(model.state_dict(), \"best_model.pt\")\n        print(\"  → New best model saved\")\n    else:\n        patience_counter += 1\n        print(f\"  → No improvement ({patience_counter}/{PATIENCE})\")\n\n        if patience_counter >= PATIENCE:\n            print(\"Early stopping triggered.\")\n            break\n\n# --------------------\n# Load best model\n# --------------------\nmodel.load_state_dict(torch.load(\"best_model.pt\"))\nprint(\"Best model loaded for inference.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T08:57:05.447903Z","iopub.execute_input":"2026-02-26T08:57:05.448424Z","iopub.status.idle":"2026-02-26T09:09:08.256638Z","shell.execute_reply.started":"2026-02-26T08:57:05.448387Z","shell.execute_reply":"2026-02-26T09:09:08.255655Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7. Sliding Window + 5 Stochastic Predictions","metadata":{}},{"cell_type":"code","source":"\ndef sliding_predict(model, seq, stochastic=False, stride=None):\n\n    if stride is None:\n        stride = MAX_LEN // 2\n\n    seq_len = len(seq)\n    preds = np.zeros((seq_len, 3), dtype=np.float32)\n    counts = np.zeros(seq_len, dtype=np.float32)\n\n    if stochastic:\n        model.train()   # Enable dropout\n    else:\n        model.eval()\n\n    with torch.no_grad():\n\n        for start in range(0, seq_len, stride):\n\n            end = min(start + MAX_LEN, seq_len)\n            chunk = seq[start:end]\n\n            tokens = np.array([NUC_MAP[s] for s in chunk], dtype=np.int64)\n            L = len(tokens)\n\n            if L < MAX_LEN:\n                tokens = np.pad(tokens, (0, MAX_LEN-L), constant_values=PAD_TOKEN)\n\n            tokens = torch.tensor(tokens).unsqueeze(0).to(DEVICE)\n            mask = torch.zeros_like(tokens).float()\n            mask[:, :L] = 1.0\n\n            pred = model(tokens, mask).squeeze(0)[:L]\n\n            # 🔥 Add small stochastic coordinate noise\n            if stochastic:\n                noise = torch.randn_like(pred) * 0.01\n                pred = pred + noise\n\n            pred = pred.cpu().numpy()\n\n            preds[start:end] += pred\n            counts[start:end] += 1\n\n    counts[counts == 0] = 1\n    return preds / counts[:, None]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T08:40:26.148356Z","iopub.execute_input":"2026-02-26T08:40:26.14908Z","iopub.status.idle":"2026-02-26T08:40:26.156636Z","shell.execute_reply.started":"2026-02-26T08:40:26.149047Z","shell.execute_reply":"2026-02-26T08:40:26.155925Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 8. Submission (5 Structures per Target)","metadata":{}},{"cell_type":"code","source":"sample = pd.read_csv(os.path.join(ROOT, \"sample_submission.csv\"))\nsubmission = sample.copy()\n\nxyz_cols = [col for col in submission.columns if col.startswith(\"x_\") \n            or col.startswith(\"y_\") \n            or col.startswith(\"z_\")]\n\nsubmission[xyz_cols] = submission[xyz_cols].astype(np.float32)\n\nfor _, row in test_seq.iterrows():\n\n    tid = row[\"target_id\"]\n    seq = row[\"sequence\"]\n\n    # Generate 5 stochastic predictions\n    preds = []\n    for i in range(5):\n        torch.manual_seed(torch.randint(0, 1_000_000, (1,)).item())\n        pred = sliding_predict(model, seq, stochastic=True)\n        pred = pred * 50.0\n        preds.append(pred)\n\n    for i in range(len(seq)):\n\n        row_id = f\"{tid}_{i+1}\"\n        idx = submission[\"ID\"] == row_id\n\n        for k in range(5):\n            submission.loc[idx, f\"x_{k+1}\"] = float(preds[k][i, 0])\n            submission.loc[idx, f\"y_{k+1}\"] = float(preds[k][i, 1])\n            submission.loc[idx, f\"z_{k+1}\"] = float(preds[k][i, 2])\n\nsubmission.to_csv(\"submission.csv\", index=False)\n\nprint(\"Submission saved with 5 structures per target.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T08:40:32.5273Z","iopub.execute_input":"2026-02-26T08:40:32.52782Z","iopub.status.idle":"2026-02-26T08:41:11.803708Z","shell.execute_reply.started":"2026-02-26T08:40:32.527791Z","shell.execute_reply":"2026-02-26T08:41:11.803058Z"}},"outputs":[],"execution_count":null}]}