{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","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":[{"sourceType":"competition","sourceId":118765,"databundleVersionId":15231210}],"dockerImageVersionId":31259,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# RNA 3D Folding – v7.5 COMPLETE Baseline\n\nIncludes:\n- Stable dataset\n- Transformer model\n- Train/validation split\n- Inference\n- Submission file creation\n","metadata":{}},{"cell_type":"code","source":"# =========================\n# Imports\n# =========================\n\nimport os\nimport random\nimport torch\nimport torch.nn as nn\nimport numpy as np\nimport pandas as pd\n\nfrom torch.utils.data import Dataset, DataLoader\n\n# =========================\n# Global Configuration\n# =========================\n\nMAX_LEN = 512\n\n# Vocabulary\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\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T18:45:31.478092Z","iopub.execute_input":"2026-02-22T18:45:31.478364Z","iopub.status.idle":"2026-02-22T18:45:35.483761Z","shell.execute_reply.started":"2026-02-22T18:45:31.47834Z","shell.execute_reply":"2026-02-22T18:45:35.482917Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Load Data (Kaggle paths)","metadata":{}},{"cell_type":"code","source":"# =========================\n# Data Loading\n# =========================\n\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\n# =========================\n# Coordinate dtype cleanup\n# =========================\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\n# =========================\n# Extract target_id from ID\n# =========================\n\ntrain_lab[\"target_id\"] = train_lab[\"ID\"].str.split(\"_\").str[0]\nval_lab[\"target_id\"]   = val_lab[\"ID\"].str.split(\"_\").str[0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T18:45:35.485172Z","iopub.execute_input":"2026-02-22T18:45:35.485643Z","iopub.status.idle":"2026-02-22T18:45:56.572881Z","shell.execute_reply.started":"2026-02-22T18:45:35.485607Z","shell.execute_reply":"2026-02-22T18:45:56.572138Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Max train length:\", train_seq[\"sequence\"].str.len().max())\nprint(\"Max val length:\", val_seq[\"sequence\"].str.len().max())\nprint(\"Max test length:\", test_seq[\"sequence\"].str.len().max())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T18:45:56.573621Z","iopub.execute_input":"2026-02-22T18:45:56.573919Z","iopub.status.idle":"2026-02-22T18:45:56.582361Z","shell.execute_reply.started":"2026-02-22T18:45:56.573896Z","shell.execute_reply":"2026-02-22T18:45:56.581361Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Train targets:\", train_seq[\"target_id\"].nunique())\nprint(\"Train labels targets:\", train_lab[\"target_id\"].nunique())\n\nprint(\"Val targets:\", val_seq[\"target_id\"].nunique())\nprint(\"Val labels targets:\", val_lab[\"target_id\"].nunique())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T18:45:56.584178Z","iopub.execute_input":"2026-02-22T18:45:56.584435Z","iopub.status.idle":"2026-02-22T18:45:57.068913Z","shell.execute_reply.started":"2026-02-22T18:45:56.584413Z","shell.execute_reply":"2026-02-22T18:45:57.068278Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(train_lab.columns)\ntrain_lab.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T18:45:57.069818Z","iopub.execute_input":"2026-02-22T18:45:57.070147Z","iopub.status.idle":"2026-02-22T18:45:57.09347Z","shell.execute_reply.started":"2026-02-22T18:45:57.070116Z","shell.execute_reply":"2026-02-22T18:45:57.092872Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"# =========================\n# RNA Dataset\n# =========================\n\nclass RNADataset(Dataset):\n\n    def __init__(self, seq_df, lab_df=None):\n        self.samples = []\n\n        if lab_df is not None:\n            lab_df = lab_df.copy()\n\n        seq_map = dict(zip(seq_df[\"target_id\"], seq_df[\"sequence\"]))\n\n        if lab_df is not None:\n\n            lab_df[[\"x_1\", \"y_1\", \"z_1\"]] = lab_df[[\"x_1\", \"y_1\", \"z_1\"]].astype(np.float32)\n\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                    # Remove NaN / Inf\n                    good = np.isfinite(xyz).all(axis=1)\n                    \n                    # Remove absurdly large values (corrupted structures)\n                    good = good & (np.abs(xyz).max(axis=1) < 1e4)\n                    \n                    r = r[good]\n                    xyz = xyz[good]\n\n\n                    # Bounds check\n                    ok = (r >= 0) & (r < L_copy)\n                    r = r[ok]\n                    xyz = xyz[ok]\n\n                    coords[r] = xyz\n                    mask[r] = 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((tid, seq_slice, coords, mask))\n\n        else:\n            for tid, seq in seq_map.items():\n                self.samples.append((tid, seq, None, None))\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n\n        tid, seq, coords, mask = self.samples[idx]\n\n        # Convert nucleotides to tokens\n        tokens = np.array([NUC_MAP[s] for s in seq], dtype=np.int64)\n\n        # Truncate\n        tokens = tokens[:MAX_LEN]\n        L = len(tokens)\n\n        # 🔹 Correct PAD handling\n        if L < MAX_LEN:\n            pad_len = MAX_LEN - L\n            tokens = np.pad(tokens, (0, pad_len), constant_values=PAD_TOKEN)\n\n        tokens = torch.tensor(tokens, dtype=torch.long)\n\n        if coords is None:\n            return tid, tokens\n\n        coords = coords[:MAX_LEN]\n        mask   = mask[:MAX_LEN]\n\n        if L < MAX_LEN:\n            pad_len = MAX_LEN - L\n            coords = np.pad(coords, ((0, pad_len), (0, 0)))\n            mask   = np.pad(mask, (0, pad_len))\n\n        # Center structure (translation invariance)\n        valid = mask == 1\n        if valid.sum() > 0:\n            center = coords[valid].mean(axis=0, keepdims=True)\n            coords[valid] -= center\n\n        # Scale coordinates\n        coords = coords / 50.0\n\n        coords = torch.tensor(coords, dtype=torch.float32)\n        mask   = torch.tensor(mask, dtype=torch.float32)\n\n        return tokens, mask, coords\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T18:45:57.09427Z","iopub.execute_input":"2026-02-22T18:45:57.094458Z","iopub.status.idle":"2026-02-22T18:45:57.107895Z","shell.execute_reply.started":"2026-02-22T18:45:57.094432Z","shell.execute_reply":"2026-02-22T18:45:57.107188Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Train/Validation Split","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=8, shuffle=True)\nval_loader   = DataLoader(val_dataset, batch_size=8, shuffle=False)\n\nprint(\"Train dataset:\", len(train_dataset))\nprint(\"Val dataset:\", len(val_dataset))\n\ntokens, mask, coords = train_dataset[0]\nprint(tokens.shape, coords.shape, mask.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T18:45:57.108817Z","iopub.execute_input":"2026-02-22T18:45:57.109163Z","iopub.status.idle":"2026-02-22T18:46:11.039941Z","shell.execute_reply.started":"2026-02-22T18:45:57.109133Z","shell.execute_reply":"2026-02-22T18:46:11.039206Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Pak één sample uit je dataset\ntokens, mask, coords = train_dataset[0]\n\n# Alleen geldige residuen (geen padding)\nvalid = mask == 1\ncoords_valid = coords[valid]\n\nprint(\"Mean absolute coord value:\", coords_valid.abs().mean().item())\nprint(\"Max absolute coord value:\", coords_valid.abs().max().item())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T18:46:11.040858Z","iopub.execute_input":"2026-02-22T18:46:11.041093Z","iopub.status.idle":"2026-02-22T18:46:11.077179Z","shell.execute_reply.started":"2026-02-22T18:46:11.041072Z","shell.execute_reply":"2026-02-22T18:46:11.076583Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"# =========================\n# Transformer Model\n# =========================\n\nclass RNATransformer(nn.Module):\n\n    def __init__(self, d_model=128, nhead=8, num_layers=4):\n        super().__init__()\n\n        # Embedding now supports PAD token\n        self.embedding = nn.Embedding(VOCAB_SIZE, d_model)\n\n        self.pos_embedding = nn.Parameter(\n            torch.zeros(1, MAX_LEN, d_model)\n        )\n\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=d_model,\n            nhead=nhead,\n            dim_feedforward=512,\n            batch_first=True\n        )\n\n        self.encoder = nn.TransformerEncoder(encoder_layer, num_layers)\n\n        self.head = nn.Linear(d_model, 3)\n\n    def forward(self, tokens, mask):\n\n        x = self.embedding(tokens)\n\n        x = x + self.pos_embedding[:, :x.size(1), :]\n\n        # True padding mask (PAD positions only)\n        padding_mask = (mask == 0)\n\n        x = self.encoder(x, src_key_padding_mask=padding_mask)\n\n        return self.head(x)\n\n\n# =========================\n# Pairwise Distance Utility\n# =========================\n\ndef pairwise_distances(coords, mask):\n    \"\"\"\n    coords: [B, L, 3]\n    mask:   [B, L]\n    \"\"\"\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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T18:46:11.077875Z","iopub.execute_input":"2026-02-22T18:46:11.078103Z","iopub.status.idle":"2026-02-22T18:46:11.085534Z","shell.execute_reply.started":"2026-02-22T18:46:11.078078Z","shell.execute_reply":"2026-02-22T18:46:11.084855Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"import torch.nn.functional as F\n\nmodel = RNATransformer().to(DEVICE)\noptimizer = torch.optim.AdamW(model.parameters(), lr=3e-5)\n\nEPOCHS = 5\n\nfor epoch in range(EPOCHS):\n\n    # =========================\n    # TRAIN\n    # =========================\n    model.train()\n    train_total = 0.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        # Pairwise distances\n        pred_dist, valid = pairwise_distances(pred, mask)\n        true_dist, _     = pairwise_distances(coords, mask)\n\n        # SmoothL1 distance loss\n        diff = F.smooth_l1_loss(pred_dist, true_dist, reduction=\"none\")\n        diff = diff * valid\n\n        denom = valid.sum()\n\n        if denom == 0:\n            continue\n\n        loss = diff.sum() / denom\n\n        # Safety checks\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.0\n    val_batches = 0\n\n    with torch.no_grad():\n\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\n    val_loss = val_total / max(val_batches, 1)\n\n    print(f\"Epoch {epoch+1} | Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T18:54:58.94069Z","iopub.execute_input":"2026-02-22T18:54:58.941087Z","iopub.status.idle":"2026-02-22T18:57:22.604061Z","shell.execute_reply.started":"2026-02-22T18:54:58.941055Z","shell.execute_reply":"2026-02-22T18:57:22.603143Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for tokens, mask, coords in val_loader:\n    print(mask.sum())\n    break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T18:59:15.537504Z","iopub.execute_input":"2026-02-22T18:59:15.538223Z","iopub.status.idle":"2026-02-22T18:59:15.547623Z","shell.execute_reply.started":"2026-02-22T18:59:15.538194Z","shell.execute_reply":"2026-02-22T18:59:15.546854Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Inference & Submission","metadata":{}},{"cell_type":"code","source":"model.eval()\n\nsample = pd.read_csv(os.path.join(ROOT, \"sample_submission.csv\"))\nsubmission = sample.copy()\n\n# Zet alle xyz kolommen naar float\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\nwith torch.no_grad():\n\n    for _, row in test_seq.iterrows():\n\n        tid = row[\"target_id\"]\n        seq = row[\"sequence\"]\n        seq_len = len(seq)\n\n        for start in range(0, seq_len, MAX_LEN):\n\n            end = min(start + MAX_LEN, seq_len)\n            seq_chunk = seq[start:end]\n\n            tokens = np.array([NUC_MAP[s] for s in seq_chunk], dtype=np.int64)\n\n            L = len(tokens)\n            if L < MAX_LEN:\n                pad = MAX_LEN - L\n                tokens = np.pad(tokens, (0, pad))\n\n            tokens = torch.tensor(tokens, dtype=torch.long).unsqueeze(0).to(DEVICE)\n            mask = torch.ones_like(tokens).float().to(DEVICE)\n\n            pred = model(tokens, mask)\n            pred = pred * 50.0\n            pred = pred.squeeze(0).cpu().numpy()\n\n            for i in range(end - start):\n\n                resid_global = start + i + 1\n                row_id = f\"{tid}_{resid_global}\"\n\n                idx = submission[\"ID\"] == row_id\n\n                submission.loc[idx, \"x_1\"] = float(pred[i, 0])\n                submission.loc[idx, \"y_1\"] = float(pred[i, 1])\n                submission.loc[idx, \"z_1\"] = float(pred[i, 2])\n\nsubmission.to_csv(\"submission.csv\", index=False)\n\nprint(\"Submission shape:\", submission.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T18:59:18.706301Z","iopub.execute_input":"2026-02-22T18:59:18.707078Z","iopub.status.idle":"2026-02-22T18:59:33.03075Z","shell.execute_reply.started":"2026-02-22T18:59:18.707041Z","shell.execute_reply":"2026-02-22T18:59:33.030023Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission.to_csv(\"submission.csv\", index=False)\n\n# ---- CHECKS ----\n\nprint(\"Submission rows:\", len(submission))\n\nexpected = 0\nfor _, row in test_seq.iterrows():\n    expected += len(row[\"sequence\"])\n\nprint(\"Expected rows:\", expected)\n\nprint(\"\\nNaNs per column:\")\nprint(submission.isna().sum())\n\nprint(\"\\nColumns:\")\nprint(submission.columns)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T18:48:47.621851Z","iopub.execute_input":"2026-02-22T18:48:47.622154Z","iopub.status.idle":"2026-02-22T18:48:47.710438Z","shell.execute_reply.started":"2026-02-22T18:48:47.622129Z","shell.execute_reply":"2026-02-22T18:48:47.709854Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample = pd.read_csv(os.path.join(ROOT, \"sample_submission.csv\"))\nprint(sample.head())\nprint(sample.columns)\nprint(len(sample))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T18:48:47.711358Z","iopub.execute_input":"2026-02-22T18:48:47.711905Z","iopub.status.idle":"2026-02-22T18:48:47.734903Z","shell.execute_reply.started":"2026-02-22T18:48:47.711872Z","shell.execute_reply":"2026-02-22T18:48:47.734269Z"}},"outputs":[],"execution_count":null}]}